204 lines
7.3 KiB
Python
204 lines
7.3 KiB
Python
from __future__ import annotations
|
|
|
|
import time
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
from typing import Any
|
|
|
|
from .antigravity_client import AntigravityClient
|
|
from .cloudcode import load_or_onboard_project
|
|
from .credentials import CredentialStore, load_agy_keychain_credentials
|
|
from .errors import ProxyError, TokenExpired
|
|
from .oauth import refresh_access_token
|
|
from .openai_compat import ChatRequest, parse_chat_request, to_openai_completion
|
|
from .transform import build_generate_content_request
|
|
|
|
|
|
def load_antigravity_credentials(store: Any | None = None) -> dict[str, Any]:
|
|
"""Load, refresh, and persist credentials for plugin use."""
|
|
if store is None:
|
|
store = CredentialStore.default()
|
|
keychain = load_agy_keychain_credentials()
|
|
from_keychain = bool(keychain)
|
|
stored = {} if from_keychain else store.load()
|
|
creds = {**stored, **keychain}
|
|
dirty = False
|
|
|
|
access = creds.get("access_token") or creds.get("access") or creds.get("token")
|
|
refresh = creds.get("refresh_token") or creds.get("refresh")
|
|
project = creds.get("project_id") or creds.get("projectId")
|
|
if not access and not refresh:
|
|
raise ProxyError(
|
|
"Missing Antigravity credentials. Run `hermes agy login`.",
|
|
status=401,
|
|
error_type="invalid_request_error",
|
|
)
|
|
expires = creds.get("expires_at") or creds.get("expires")
|
|
if refresh and (not access or (isinstance(expires, (int, float)) and time.time() + 60 >= float(expires))):
|
|
refreshed = refresh_access_token(str(refresh))
|
|
creds.update(refreshed)
|
|
access = refreshed["access_token"]
|
|
refresh = creds.get("refresh_token") or creds.get("refresh")
|
|
dirty = not from_keychain
|
|
if not project:
|
|
if not access:
|
|
raise ProxyError(
|
|
"Missing access token for Antigravity project discovery",
|
|
status=401,
|
|
error_type="invalid_request_error",
|
|
)
|
|
project = load_or_onboard_project(str(access))
|
|
creds["project_id"] = project
|
|
dirty = not from_keychain
|
|
if dirty:
|
|
store.save(creds)
|
|
return {
|
|
"access_token": str(access or creds["access_token"]),
|
|
"refresh_token": str(refresh or creds.get("refresh_token") or ""),
|
|
"project_id": str(project),
|
|
"source": "agy-keychain" if from_keychain else "store",
|
|
}
|
|
|
|
|
|
def build_upstream_body(request: ChatRequest, *, store: Any | None = None) -> tuple[dict[str, Any], dict[str, Any]]:
|
|
creds = load_antigravity_credentials(store)
|
|
body = build_generate_content_request(
|
|
model=request.model,
|
|
project_id=creds["project_id"],
|
|
messages=request.messages,
|
|
tools=request.tools,
|
|
reasoning_effort=request.reasoning_effort,
|
|
max_tokens=request.max_tokens,
|
|
temperature=request.temperature,
|
|
top_p=request.top_p,
|
|
tool_choice=request.tool_choice,
|
|
)
|
|
return body, creds
|
|
|
|
|
|
def generate_chat_completion(
|
|
payload: dict[str, Any],
|
|
*,
|
|
client: Any | None = None,
|
|
store: Any | None = None,
|
|
) -> dict[str, Any]:
|
|
"""Execute one OpenAI-shaped chat request against Antigravity in-process."""
|
|
request = parse_chat_request(payload)
|
|
if client is None:
|
|
client = AntigravityClient()
|
|
if store is None:
|
|
store = CredentialStore.default()
|
|
body, creds = build_upstream_body(request, store=store)
|
|
try:
|
|
upstream = client.generate(access_token=creds["access_token"], body=body)
|
|
except TokenExpired:
|
|
if not creds.get("refresh_token"):
|
|
raise
|
|
refreshed = refresh_access_token(creds["refresh_token"])
|
|
if creds.get("source") == "store":
|
|
saved = store.load()
|
|
saved.update(refreshed)
|
|
store.save(saved)
|
|
upstream = client.generate(access_token=refreshed["access_token"], body=body)
|
|
return to_openai_completion(request.model, upstream)
|
|
|
|
|
|
def _namespace(value: Any) -> Any:
|
|
if isinstance(value, dict):
|
|
return SimpleNamespace(**{k: _namespace(v) for k, v in value.items()})
|
|
if isinstance(value, list):
|
|
return [_namespace(v) for v in value]
|
|
return value
|
|
|
|
|
|
def format_antigravity_error(err: Any) -> str:
|
|
"""Format an error message with 'Antigravity error: ' prefix without duplicate prefixes."""
|
|
msg = err.get("message") if isinstance(err, dict) else str(err or "unknown error")
|
|
msg = msg.strip()
|
|
|
|
prefixes_to_strip = [
|
|
"Antigravity error:",
|
|
"Antigravity (agy) error:",
|
|
"agy error:",
|
|
"Antigravity error",
|
|
"agy error",
|
|
]
|
|
changed = True
|
|
while changed:
|
|
changed = False
|
|
for p in prefixes_to_strip:
|
|
if msg.lower().startswith(p.lower()):
|
|
msg = msg[len(p):].strip(" :")
|
|
changed = True
|
|
break
|
|
|
|
return f"Antigravity error: {msg}" if msg else "Antigravity error: unknown error"
|
|
|
|
|
|
def openai_completion_object(completion: dict[str, Any]) -> SimpleNamespace:
|
|
"""Return an object compatible with Hermes' ChatCompletionsTransport."""
|
|
completion = dict(completion)
|
|
choices = []
|
|
|
|
# Handle error payload without crashing choices[0]
|
|
if "error" in completion and not completion.get("choices"):
|
|
err_text = format_antigravity_error(completion.get("error"))
|
|
choices.append({
|
|
"index": 0,
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": err_text,
|
|
"tool_calls": None,
|
|
},
|
|
"finish_reason": "stop",
|
|
})
|
|
else:
|
|
for raw_choice in completion.get("choices") or []:
|
|
choice = dict(raw_choice)
|
|
message = dict(choice.get("message") or {})
|
|
message.setdefault("content", None)
|
|
message.setdefault("tool_calls", None)
|
|
choice["message"] = message
|
|
choices.append(choice)
|
|
|
|
# Fallback to ensure choices is never empty
|
|
if not choices:
|
|
choices.append({
|
|
"index": 0,
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": "Empty or unrecognized completion response from provider.",
|
|
"tool_calls": None,
|
|
},
|
|
"finish_reason": "stop",
|
|
})
|
|
|
|
completion["choices"] = choices
|
|
completion.setdefault("usage", {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0})
|
|
return _namespace(completion)
|
|
|
|
|
|
def ensure_provider_profile_files(root: Path | None = None) -> Path:
|
|
"""Install the tiny model-provider profile that makes `hermes model` see Antigravity."""
|
|
if root is None:
|
|
try:
|
|
from hermes_constants import get_hermes_home
|
|
|
|
root = get_hermes_home()
|
|
except Exception:
|
|
root = Path.home() / ".hermes"
|
|
plugin_dir = Path(root).expanduser() / "plugins" / "model-providers" / "antigravity"
|
|
plugin_dir.mkdir(parents=True, exist_ok=True)
|
|
(plugin_dir / "__init__.py").write_text(
|
|
"from antigravity_provider.hermes_provider import register_provider_profile\n"
|
|
"register_provider_profile()\n",
|
|
encoding="utf-8",
|
|
)
|
|
(plugin_dir / "plugin.yaml").write_text(
|
|
"name: antigravity\n"
|
|
"kind: model-provider\n"
|
|
"version: 0.1.0\n"
|
|
"description: Google Antigravity provider profile\n",
|
|
encoding="utf-8",
|
|
)
|
|
return plugin_dir
|