"""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
