rebase
This commit is contained in:
462
app/core/llm.py
Normal file
462
app/core/llm.py
Normal file
@@ -0,0 +1,462 @@
|
||||
"""LLM client — OpenAI-compatible API wrapper with retry, logging, and streaming.
|
||||
|
||||
Usage:
|
||||
client = LlmClient.from_settings(settings_dict)
|
||||
resp = await client.complete(
|
||||
stage="orchestrator_phase1",
|
||||
messages=[{"role": "system", "content": "..."}, ...],
|
||||
tools=[...], # optional
|
||||
temperature=0.7,
|
||||
max_tokens=2048,
|
||||
stream=False, # if True, returns an async iterator of deltas
|
||||
user_id=..., world_id=..., step_id=...,
|
||||
)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.logging import get_logger
|
||||
from app.models import LlmCallLog
|
||||
|
||||
_logger = get_logger(__name__)
|
||||
|
||||
|
||||
class LLMError(Exception):
|
||||
"""Base LLM error."""
|
||||
|
||||
def __init__(self, code: str, message: str, status: str = "api_error"):
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
self.status = status
|
||||
|
||||
|
||||
class LLMTimeoutError(LLMError):
|
||||
def __init__(self, message: str = "LLM call timed out"):
|
||||
super().__init__("llm_timeout", message, status="timeout")
|
||||
|
||||
|
||||
class LLMUnavailableError(LLMError):
|
||||
def __init__(self, message: str = "LLM provider unavailable"):
|
||||
super().__init__("llm_unavailable", message, status="api_error")
|
||||
|
||||
|
||||
class LLMResponseError(LLMError):
|
||||
def __init__(self, message: str, code: str = "parse_error"):
|
||||
super().__init__(code, message, status="parse_error")
|
||||
|
||||
|
||||
class LlmClient:
|
||||
"""OpenAI-compatible LLM client with retry, logging, and streaming."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_url: str,
|
||||
api_key: str,
|
||||
model: str,
|
||||
timeout: float = 60.0,
|
||||
max_retries: int = 3,
|
||||
):
|
||||
self._api_url = api_url.rstrip("/")
|
||||
self._api_key = api_key
|
||||
self._model = model
|
||||
self._timeout = timeout
|
||||
self._max_retries = max_retries
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Construction
|
||||
# ------------------------------------------------------------------ #
|
||||
@classmethod
|
||||
def from_settings(cls, settings: dict[str, Any]) -> "LlmClient":
|
||||
return cls(
|
||||
api_url=settings.get("llm.api_url", "http://localhost:11434/v1"),
|
||||
api_key=settings.get("llm.api_key", ""),
|
||||
model=settings.get("llm.model", "qwen2.5-7b-instruct"),
|
||||
timeout=float(settings.get("llm.timeout_seconds", 60)),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Non-streaming call
|
||||
# ------------------------------------------------------------------ #
|
||||
async def complete(
|
||||
self,
|
||||
*,
|
||||
stage: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict] | None = None,
|
||||
tool_choice: Any = None,
|
||||
temperature: float = 0.7,
|
||||
top_p: float = 0.9,
|
||||
max_tokens: int = 2048,
|
||||
user_id: uuid.UUID | None = None,
|
||||
world_id: uuid.UUID | None = None,
|
||||
step_id: uuid.UUID | None = None,
|
||||
session: AsyncSession | None = None,
|
||||
stream: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""Make a non-streaming chat completion call.
|
||||
|
||||
Returns a dict with keys:
|
||||
- `message`: assistant message (with `content` and optional `tool_calls`)
|
||||
- `finish_reason`: stop | length | tool_calls
|
||||
- `prompt_tokens`, `completion_tokens`, `latency_ms`
|
||||
- `log_id`: id of the LlmCallLog row if `session` provided
|
||||
"""
|
||||
if stream:
|
||||
raise ValueError("Use stream_complete() for streaming calls")
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": self._model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"top_p": top_p,
|
||||
"max_tokens": max_tokens,
|
||||
"stream": False,
|
||||
}
|
||||
if tools:
|
||||
payload["tools"] = tools
|
||||
payload["tool_choice"] = tool_choice or "auto"
|
||||
|
||||
start = time.monotonic()
|
||||
last_exc: Exception | None = None
|
||||
for attempt in range(self._max_retries):
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as client:
|
||||
resp = await client.post(
|
||||
f"{self._api_url}/chat/completions",
|
||||
headers={
|
||||
"Authorization": f"Bearer {self._api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
json=payload,
|
||||
)
|
||||
if resp.status_code >= 500:
|
||||
raise LLMUnavailableError(
|
||||
f"LLM provider returned {resp.status_code}: {resp.text[:200]}"
|
||||
)
|
||||
if resp.status_code == 429:
|
||||
raise LLMUnavailableError("LLM provider rate-limited (429)")
|
||||
if resp.status_code >= 400:
|
||||
raise LLMResponseError(
|
||||
f"LLM provider returned {resp.status_code}: {resp.text[:500]}",
|
||||
code="api_error",
|
||||
)
|
||||
data = resp.json()
|
||||
break
|
||||
except (httpx.TimeoutException, asyncio.TimeoutError) as e:
|
||||
last_exc = LLMTimeoutError(str(e))
|
||||
_logger.warning(
|
||||
"llm_timeout", stage=stage, attempt=attempt + 1, error=str(e)
|
||||
)
|
||||
except (httpx.ConnectError, httpx.NetworkError) as e:
|
||||
last_exc = LLMUnavailableError(str(e))
|
||||
_logger.warning(
|
||||
"llm_connection_error", stage=stage, attempt=attempt + 1, error=str(e)
|
||||
)
|
||||
except LLMError as e:
|
||||
last_exc = e
|
||||
_logger.warning(
|
||||
"llm_error", stage=stage, attempt=attempt + 1, error=str(e)
|
||||
)
|
||||
# exponential backoff
|
||||
await asyncio.sleep(min(2**attempt, 4))
|
||||
else:
|
||||
# All retries exhausted
|
||||
if session is not None:
|
||||
await self._write_log_safely(
|
||||
session=session,
|
||||
stage=stage,
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
response_message={},
|
||||
tool_calls=None,
|
||||
prompt_tokens=None,
|
||||
completion_tokens=None,
|
||||
latency_ms=int((time.monotonic() - start) * 1000),
|
||||
temperature=temperature,
|
||||
status=last_exc.status if isinstance(last_exc, LLMError) else "api_error",
|
||||
error_message=str(last_exc) if last_exc else "unknown",
|
||||
user_id=user_id,
|
||||
world_id=world_id,
|
||||
step_id=step_id,
|
||||
)
|
||||
assert last_exc is not None
|
||||
raise last_exc
|
||||
|
||||
latency_ms = int((time.monotonic() - start) * 1000)
|
||||
choice = data["choices"][0]
|
||||
msg = choice.get("message", {})
|
||||
finish_reason = choice.get("finish_reason", "stop")
|
||||
usage = data.get("usage", {})
|
||||
|
||||
log_id: uuid.UUID | None = None
|
||||
if session is not None:
|
||||
log_id = await self._write_log_safely(
|
||||
session=session,
|
||||
stage=stage,
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
response_message=msg,
|
||||
tool_calls=msg.get("tool_calls"),
|
||||
prompt_tokens=usage.get("prompt_tokens"),
|
||||
completion_tokens=usage.get("completion_tokens"),
|
||||
latency_ms=latency_ms,
|
||||
temperature=temperature,
|
||||
status="ok",
|
||||
error_message=None,
|
||||
user_id=user_id,
|
||||
world_id=world_id,
|
||||
step_id=step_id,
|
||||
)
|
||||
|
||||
return {
|
||||
"message": msg,
|
||||
"finish_reason": finish_reason,
|
||||
"prompt_tokens": usage.get("prompt_tokens"),
|
||||
"completion_tokens": usage.get("completion_tokens"),
|
||||
"latency_ms": latency_ms,
|
||||
"log_id": log_id,
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Streaming call
|
||||
# ------------------------------------------------------------------ #
|
||||
async def stream_complete(
|
||||
self,
|
||||
*,
|
||||
stage: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict] | None = None,
|
||||
tool_choice: Any = None,
|
||||
temperature: float = 0.85,
|
||||
top_p: float = 0.95,
|
||||
max_tokens: int = 2048,
|
||||
user_id: uuid.UUID | None = None,
|
||||
world_id: uuid.UUID | None = None,
|
||||
step_id: uuid.UUID | None = None,
|
||||
session: AsyncSession | None = None,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
"""Stream chat completion. Yields dicts with keys:
|
||||
- `delta`: {content?, tool_calls?}
|
||||
- `finish_reason`: present only on the final chunk
|
||||
After the iterator is exhausted, the call is logged to `llm_call_logs`.
|
||||
"""
|
||||
payload: dict[str, Any] = {
|
||||
"model": self._model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"top_p": top_p,
|
||||
"max_tokens": max_tokens,
|
||||
"stream": True,
|
||||
}
|
||||
if tools:
|
||||
payload["tools"] = tools
|
||||
payload["tool_choice"] = tool_choice or "auto"
|
||||
|
||||
start = time.monotonic()
|
||||
full_content_parts: list[str] = []
|
||||
full_tool_calls: list[dict] = []
|
||||
finish_reason: str | None = None
|
||||
usage: dict[str, Any] = {}
|
||||
status = "ok"
|
||||
error_message: str | None = None
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as client:
|
||||
async with client.stream(
|
||||
"POST",
|
||||
f"{self._api_url}/chat/completions",
|
||||
headers={
|
||||
"Authorization": f"Bearer {self._api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
json=payload,
|
||||
) as resp:
|
||||
if resp.status_code >= 400:
|
||||
body = await resp.aread()
|
||||
raise LLMResponseError(
|
||||
f"LLM provider returned {resp.status_code}: {body.decode('utf-8', 'ignore')[:500]}",
|
||||
code="api_error",
|
||||
)
|
||||
async for line in resp.aiter_lines():
|
||||
if not line:
|
||||
continue
|
||||
if line.startswith("data: "):
|
||||
line = line[6:]
|
||||
if line.strip() == "[DONE]":
|
||||
break
|
||||
try:
|
||||
chunk = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if not chunk.get("choices"):
|
||||
if chunk.get("usage"):
|
||||
usage = chunk["usage"]
|
||||
continue
|
||||
choice = chunk["choices"][0]
|
||||
delta = choice.get("delta", {})
|
||||
if delta.get("content"):
|
||||
full_content_parts.append(delta["content"])
|
||||
if delta.get("tool_calls"):
|
||||
full_tool_calls.extend(delta["tool_calls"])
|
||||
if choice.get("finish_reason"):
|
||||
finish_reason = choice["finish_reason"]
|
||||
yield {"delta": delta, "finish_reason": finish_reason}
|
||||
except Exception as e:
|
||||
status = "api_error" if not isinstance(e, LLMTimeoutError) else "timeout"
|
||||
error_message = str(e)
|
||||
_logger.warning("llm_stream_error", stage=stage, error=error_message)
|
||||
raise
|
||||
finally:
|
||||
latency_ms = int((time.monotonic() - start) * 1000)
|
||||
if session is not None:
|
||||
full_content = "".join(full_content_parts)
|
||||
await self._write_log_safely(
|
||||
session=session,
|
||||
stage=stage,
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
response_message={
|
||||
"role": "assistant",
|
||||
"content": full_content,
|
||||
"tool_calls": full_tool_calls or None,
|
||||
},
|
||||
tool_calls=full_tool_calls or None,
|
||||
prompt_tokens=usage.get("prompt_tokens"),
|
||||
completion_tokens=usage.get("completion_tokens"),
|
||||
latency_ms=latency_ms,
|
||||
temperature=temperature,
|
||||
status=status,
|
||||
error_message=error_message,
|
||||
user_id=user_id,
|
||||
world_id=world_id,
|
||||
step_id=step_id,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Safe logging (separate transaction)
|
||||
# ------------------------------------------------------------------ #
|
||||
async def _write_log_safely(
|
||||
self,
|
||||
*,
|
||||
session: AsyncSession,
|
||||
stage: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict] | None,
|
||||
response_message: dict[str, Any],
|
||||
tool_calls: list | None,
|
||||
prompt_tokens: int | None,
|
||||
completion_tokens: int | None,
|
||||
latency_ms: int,
|
||||
temperature: float,
|
||||
status: str,
|
||||
error_message: str | None,
|
||||
user_id: uuid.UUID | None,
|
||||
world_id: uuid.UUID | None,
|
||||
step_id: uuid.UUID | None,
|
||||
) -> uuid.UUID | None:
|
||||
"""Insert an LlmCallLog row in a nested transaction so it survives rollback.
|
||||
|
||||
Errors here are logged but never raised — logging is best-effort.
|
||||
"""
|
||||
try:
|
||||
async with session.begin_nested():
|
||||
log = LlmCallLog(
|
||||
user_id=user_id,
|
||||
world_id=world_id,
|
||||
step_id=step_id,
|
||||
stage=stage,
|
||||
model=self._model,
|
||||
request_messages=messages,
|
||||
request_tools=tools,
|
||||
response_message=response_message,
|
||||
tool_calls=tool_calls,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
latency_ms=latency_ms,
|
||||
temperature=temperature,
|
||||
status=status,
|
||||
error_message=error_message,
|
||||
)
|
||||
session.add(log)
|
||||
await session.flush()
|
||||
log_id = log.id
|
||||
await session.commit()
|
||||
return log_id
|
||||
except Exception as e: # noqa: BLE001
|
||||
_logger.error("llm_log_write_failed", stage=stage, error=str(e))
|
||||
try:
|
||||
await session.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Mock LLM client (for tests)
|
||||
# --------------------------------------------------------------------------- #
|
||||
class MockLlmClient:
|
||||
"""Replay-based mock LLM client. Returns pre-recorded responses per stage."""
|
||||
|
||||
def __init__(self, replay_data: dict[str, list[dict]] | None = None):
|
||||
self._replay = replay_data or {}
|
||||
self._call_counts: dict[str, int] = {}
|
||||
# Allow recording mode
|
||||
self.recorded_calls: list[dict[str, Any]] = []
|
||||
|
||||
def set_replay(self, stage: str, responses: list[dict]) -> None:
|
||||
self._replay[stage] = responses
|
||||
self._call_counts.pop(stage, None)
|
||||
|
||||
async def complete(self, *, stage: str, messages=None, tools=None, **kwargs) -> dict[str, Any]:
|
||||
idx = self._call_counts.get(stage, 0)
|
||||
responses = self._replay.get(stage, [])
|
||||
if idx >= len(responses):
|
||||
raise LLMResponseError(
|
||||
f"Replay exhausted for stage {stage} (call #{idx + 1})",
|
||||
code="replay_exhausted",
|
||||
)
|
||||
resp = responses[idx]
|
||||
self._call_counts[stage] = idx + 1
|
||||
self.recorded_calls.append({"stage": stage, "messages": messages, "tools": tools})
|
||||
|
||||
# Mimic the real client's return shape
|
||||
return {
|
||||
"message": resp.get("message", {"role": "assistant", "content": resp.get("content", "")}),
|
||||
"finish_reason": resp.get("finish_reason", "stop"),
|
||||
"prompt_tokens": resp.get("prompt_tokens", 0),
|
||||
"completion_tokens": resp.get("completion_tokens", 0),
|
||||
"latency_ms": 0,
|
||||
"log_id": None,
|
||||
}
|
||||
|
||||
async def stream_complete(self, *, stage: str, messages=None, tools=None, **kwargs):
|
||||
idx = self._call_counts.get(stage, 0)
|
||||
responses = self._replay.get(stage, [])
|
||||
if idx >= len(responses):
|
||||
raise LLMResponseError(
|
||||
f"Replay exhausted for stage {stage} (call #{idx + 1})",
|
||||
code="replay_exhausted",
|
||||
)
|
||||
resp = responses[idx]
|
||||
self._call_counts[stage] = idx + 1
|
||||
content = resp.get("message", {}).get("content", resp.get("content", ""))
|
||||
# Yield content in 3 chunks for streaming tests
|
||||
chunk_size = max(1, len(content) // 3)
|
||||
for i in range(0, len(content), chunk_size):
|
||||
yield {"delta": {"content": content[i : i + chunk_size]}, "finish_reason": None}
|
||||
yield {"delta": {}, "finish_reason": "stop"}
|
||||
|
||||
|
||||
def get_mock_client() -> MockLlmClient:
|
||||
"""Convenience factory — used in tests and as a fallback in dev when no LLM configured."""
|
||||
return MockLlmClient()
|
||||
Reference in New Issue
Block a user