sage-api / src /llm /base.py
hiasgnpsadgd's picture
Deploy current MIMIC-only SAGE backend
de0f30b verified
Raw
History Blame Contribute Delete
2.84 kB
"""Shared contracts and JSON helpers for LLM providers."""
from __future__ import annotations
import json
from typing import Any, Protocol
class LlmProviderError(RuntimeError):
"""Raised when an LLM provider cannot return a usable response."""
class JsonLlmProvider(Protocol):
"""Provider protocol shared by real and fake JSON LLM providers."""
def complete_json(
self,
*,
system_prompt: str,
user_prompt: str,
output_schema: dict[str, Any],
metadata: dict[str, Any],
) -> dict[str, Any]:
...
def build_json_response_format(
*,
output_schema: dict[str, Any],
metadata: dict[str, Any],
use_json_schema: bool,
strict: bool,
default_schema_name: str,
) -> dict[str, Any]:
if not use_json_schema or not output_schema:
return {"type": "json_object"}
return {
"type": "json_schema",
"json_schema": {
"name": json_schema_name(metadata, default_name=default_schema_name),
"strict": strict,
"schema": output_schema,
},
}
def json_schema_name(metadata: dict[str, Any], *, default_name: str) -> str:
raw = "_".join(
str(metadata.get(key, "")).strip()
for key in ("strategy_name", "purpose", "stage", "call_id")
if metadata.get(key)
)
if not raw:
raw = default_name
safe = [character if character.isalnum() or character in {"_", "-"} else "_" for character in raw]
name = "".join(safe).strip("_") or default_name
if not name[0].isalpha():
name = f"{default_name}_{name}"
return name[:64]
def parse_json_content(content: Any) -> Any:
"""Parse model content that may wrap JSON in surrounding text."""
text = content_to_text(content).strip()
try:
return json.loads(text)
except json.JSONDecodeError:
start = text.find("{")
end = text.rfind("}")
if start >= 0 and end >= start:
return json.loads(text[start : end + 1])
raise
def extract_chat_message_content(body: dict[str, Any], *, provider_name: str) -> Any:
try:
choice = body["choices"][0]
message = choice["message"]
return message.get("content")
except (KeyError, IndexError, TypeError) as exc:
raise LlmProviderError(f"{provider_name} returned an unexpected response payload: {body!r}") from exc
def content_to_text(content: Any) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):
parts: list[str] = []
for item in content:
if isinstance(item, dict):
parts.append(str(item.get("text") or item.get("content") or ""))
else:
parts.append(str(item))
return "".join(parts)
return str(content or "")