"""Provide an OpenAI-compatible wrapper for accessing VertexAI models.""" import json import os from typing import Any import google.auth import google.auth.transport.requests import openai VERTEX_PROJECT_ID = "ml-gemini-455703" LOCATION = "us-central1" BASE_URL = f"https://{LOCATION}-aiplatform.googleapis.com/v1/projects/{VERTEX_PROJECT_ID}/locations/{LOCATION}/endpoints/openapi" def _set_up_creds(): google_creds_data = json.loads(os.environ["GOOGLE_APPLICATION_CREDENTIALS_DATA"]) google_creds_filename = "google_application_credentials.json" with open(google_creds_filename, "w") as f: json.dump(google_creds_data, f) os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = google_creds_filename class OpenAICredentialsRefresher: def __init__(self, **kwargs: Any) -> None: # Set a placeholder key here self.client = openai.OpenAI(**kwargs, api_key="PLACEHOLDER") self.creds, self.project = google.auth.default( scopes=["https://www.googleapis.com/auth/cloud-platform"] ) def __getattr__(self, name: str) -> Any: # TODO(pat) technically we could do this once every 12 hours rather than every hour if not self.creds.valid: print("refreshing creds") self.creds.refresh(google.auth.transport.requests.Request()) if not self.creds.valid: raise RuntimeError("Unable to refresh Google Vertex auth") self.client.api_key = self.creds.token return getattr(self.client, name) def make_vertex_client(): _set_up_creds() client = OpenAICredentialsRefresher(base_url=BASE_URL) return client