"""JSON LLM adapters. Offline tests use ReplayClient or ScriptedClient.""" from __future__ import annotations import json from collections import defaultdict, deque from dataclasses import dataclass, field from pathlib import Path from typing import Any, Protocol @dataclass(frozen=True) class ModelResponse: data: dict[str, Any] raw_text: str model: str response_id: str | None = None usage: dict[str, Any] = field(default_factory=dict) class JsonLLM(Protocol): model: str async def generate_json( self, *, system_prompt: str, user_prompt: str, request_id: str, ) -> ModelResponse: ... class OpenAIJsonClient: """Optional OpenAI chat-completions adapter. No API request is made until ``generate_json`` is awaited. """ def __init__( self, model: str, *, api_key: str | None = None, max_completion_tokens: int = 16_000, seed: int | None = None, ) -> None: try: from openai import AsyncOpenAI except ImportError as exc: raise RuntimeError( "Install the optional API dependency: pip install -e '.[api]'" ) from exc self.model = model self._client = AsyncOpenAI(api_key=api_key) self.max_completion_tokens = max_completion_tokens self.seed = seed async def generate_json( self, *, system_prompt: str, user_prompt: str, request_id: str, ) -> ModelResponse: kwargs: dict[str, Any] = { "model": self.model, "messages": [ {"role": "system", "content": system_prompt}, {"role": "user", "content": user_prompt}, ], "response_format": {"type": "json_object"}, "max_completion_tokens": self.max_completion_tokens, } if self.seed is not None: kwargs["seed"] = self.seed response = await self._client.chat.completions.create(**kwargs) raw = response.choices[0].message.content or "{}" data = json.loads(raw) usage = ( response.usage.model_dump() if response.usage is not None and hasattr(response.usage, "model_dump") else {} ) return ModelResponse( data=data, raw_text=raw, model=self.model, response_id=response.id, usage=usage, ) class ReplayClient: """Replay archived call responses keyed by request ID.""" def __init__(self, path: Path, model: str = "replay") -> None: self.model = model payload = json.loads(path.read_text(encoding="utf-8")) if isinstance(payload, list): self._responses = { str(row["request_id"]): row["response_data"] for row in payload } elif isinstance(payload, dict): self._responses = payload else: raise ValueError("replay file must be a JSON object or list") async def generate_json( self, *, system_prompt: str, user_prompt: str, request_id: str, ) -> ModelResponse: if request_id not in self._responses: raise KeyError(f"no replay response for {request_id}") data = self._responses[request_id] return ModelResponse( data=data, raw_text=json.dumps(data, ensure_ascii=False), model=self.model, response_id=f"replay:{request_id}", ) class ScriptedClient: """Small deterministic client for tests and offline smoke runs.""" def __init__( self, responses: dict[str, list[dict[str, Any]] | dict[str, Any]], model: str = "scripted", ) -> None: self.model = model self.calls: list[str] = [] self._responses: dict[str, deque[dict[str, Any]]] = {} for prefix, values in responses.items(): sequence = values if isinstance(values, list) else [values] self._responses[prefix] = deque(sequence) self._prefix_counts: defaultdict[str, int] = defaultdict(int) async def generate_json( self, *, system_prompt: str, user_prompt: str, request_id: str, ) -> ModelResponse: self.calls.append(request_id) prefixes = sorted( (prefix for prefix in self._responses if request_id.startswith(prefix)), key=len, reverse=True, ) if not prefixes: raise KeyError(f"no scripted response matches {request_id}") prefix = prefixes[0] queue = self._responses[prefix] if not queue: raise IndexError(f"scripted responses exhausted for {prefix}") data = queue.popleft() self._prefix_counts[prefix] += 1 return ModelResponse( data=data, raw_text=json.dumps(data, ensure_ascii=False), model=self.model, response_id=f"scripted:{prefix}:{self._prefix_counts[prefix]}", )