Files
openworker/coworker/providers/vertex_provider.py
T
Rohit C Prasad 050cc894e7 Add Google Vertex AI provider with per-family dispatch
gemini/ and claude/ ids reuse the native providers; openweight/ goes through the
MaaS OpenAI-compat endpoint with an auto-refreshed google-auth bearer.
Credentials: service-account JSON or Application Default Credentials.
2026-07-25 16:19:00 -07:00

199 lines
7.9 KiB
Python

"""Google Vertex AI provider — one entry in Settings, three wire paths by model family.
Routed ids look like `vertex:<family>/<model id>`; the router strips `vertex:` and this
provider splits the family segment, reusing an existing provider class per family:
- `gemini/…` → the native `GeminiProvider` over `genai.Client(vertexai=True)`.
- `claude/…` → the native `AnthropicProvider` over the SDK's `AnthropicVertex` client.
- `openweight/…` → `OpenAIProvider` against Vertex's OpenAI-compatible MaaS endpoint
(Llama, Qwen, DeepSeek, …; ids keep their publisher segment, e.g. `openweight/meta/…`).
An id with no recognized family segment is best-effort routed by name (gemini* → Gemini,
claude* → Claude, anything else → MaaS) so a raw id pasted without the add-model dropdown
still works.
Credentials: an explicit service-account JSON (pasted content or a file path) when the
profile has one, else Application Default Credentials (`gcloud auth application-default
login`). The MaaS path authenticates with a google-auth bearer token that expires ~hourly —
this wrapper refreshes it and rebuilds the OpenAI sub-client as needed; the two native SDK
clients take the credentials object and refresh internally.
"""
from __future__ import annotations
import json
from typing import Any, Optional
from .anthropic_provider import AnthropicProvider
from .base import AssistantTurn, ModelCapabilities, ProviderClient
from .capabilities import capabilities_for
from .gemini_provider import GeminiProvider
from .openai_provider import OpenAIProvider
_SCOPES = ["https://www.googleapis.com/auth/cloud-platform"]
_FAMILIES = ("gemini", "claude", "openweight")
def load_credentials(service_account_json: Optional[str]) -> Any:
"""Explicit service-account JSON (content or path) → Credentials; blank → None (the
SDKs and the token path then fall back to Application Default Credentials)."""
raw = (service_account_json or "").strip()
if not raw:
return None
from google.oauth2 import service_account
if raw.startswith("{"):
info = json.loads(raw)
return service_account.Credentials.from_service_account_info(
info, scopes=_SCOPES
)
return service_account.Credentials.from_service_account_file(raw, scopes=_SCOPES)
class VertexProvider(ProviderClient):
"""Family dispatcher: splits `<family>/<model id>` and delegates to the sub-client."""
def __init__(
self,
*,
project: Optional[str] = None,
location: Optional[str] = None,
service_account_json: Optional[str] = None,
credentials: Any = None,
gemini_client: Optional[ProviderClient] = None,
claude_client: Optional[ProviderClient] = None,
openweight_client: Optional[ProviderClient] = None,
):
self._project = project
self._location = location
self._service_account_json = service_account_json
self._credentials = credentials # test seam; normally resolved lazily
# Test seams: pre-built sub-providers skip the SDK construction below.
self._clients: dict[str, ProviderClient] = {}
if gemini_client is not None:
self._clients["gemini"] = gemini_client
if claude_client is not None:
self._clients["claude"] = claude_client
if openweight_client is not None:
self._clients["openweight"] = openweight_client
self._openweight_injected = openweight_client is not None
@staticmethod
def _split(model: str) -> tuple[str, str]:
if "/" in model:
family, rest = model.split("/", 1)
if family in _FAMILIES:
return family, rest
# Raw id without a family segment: route by name, best effort.
if model.startswith("gemini"):
return "gemini", model
if model.startswith("claude"):
return "claude", model
return "openweight", model
# -- credentials -------------------------------------------------------------
def _explicit_credentials(self) -> Any:
"""The service-account credentials, or None to let each SDK use ADC."""
if self._credentials is None:
self._credentials = load_credentials(self._service_account_json)
return self._credentials
def _bearer_credentials(self) -> Any:
"""Credentials for the MaaS bearer token: explicit service account, else ADC."""
creds = self._explicit_credentials()
if creds is None:
import google.auth
try:
creds, _ = google.auth.default(scopes=_SCOPES)
except Exception as exc:
raise RuntimeError(
"No Google Cloud credentials found — paste a service-account JSON "
"in Settings ▸ Models, or run `gcloud auth application-default login`."
) from exc
self._credentials = creds
return creds
# -- family sub-clients --------------------------------------------------------
def _family_client(self, family: str) -> ProviderClient:
if family == "openweight":
return self._openweight_client()
client = self._clients.get(family)
if client is None:
if family == "gemini":
from google import genai
client = GeminiProvider(
client=genai.Client(
vertexai=True,
project=self._project,
location=self._location,
credentials=self._explicit_credentials(),
)
)
else:
from anthropic import AnthropicVertex
client = AnthropicProvider(
client=AnthropicVertex(
project_id=self._project,
region=self._location,
credentials=self._explicit_credentials(),
)
)
self._clients[family] = client
return client
def _openweight_client(self) -> ProviderClient:
"""OpenAIProvider over the Vertex MaaS endpoint, rebuilt whenever the bearer
token has to be refreshed (google-auth tokens expire ~hourly)."""
if self._openweight_injected:
return self._clients["openweight"]
creds = self._bearer_credentials()
if not getattr(creds, "valid", False):
from google.auth.transport.requests import Request
creds.refresh(Request())
self._clients.pop("openweight", None) # stale token — rebuild below
client = self._clients.get("openweight")
if client is None:
base = (
f"https://{self._location}-aiplatform.googleapis.com/v1/projects/"
f"{self._project}/locations/{self._location}/endpoints/openapi"
)
client = OpenAIProvider(api_key=creds.token, base_url=base)
self._clients["openweight"] = client
return client
# -- ProviderClient -------------------------------------------------------------
def complete(
self,
*,
model: str,
messages: list[dict[str, Any]],
tools: Optional[list[dict[str, Any]]] = None,
**settings: Any,
) -> AssistantTurn:
family, rest = self._split(model)
return self._family_client(family).complete(
model=rest, messages=messages, tools=tools, **settings
)
def stream(
self,
*,
model: str,
messages: list[dict[str, Any]],
tools: Optional[list[dict[str, Any]]] = None,
**settings: Any,
):
family, rest = self._split(model)
return self._family_client(family).stream(
model=rest, messages=messages, tools=tools, **settings
)
def capabilities(self, model: str) -> ModelCapabilities:
qualified = model if model.startswith("vertex:") else f"vertex:{model}"
return capabilities_for(qualified)