"""Inference adapter for thesis rewriting via the InferenceGateway. Single implementation replacing the duplicate Ollama/vLLM branching in ``thesis_llm.py``. Uses InferenceGateway.generate() with the appropriate target — no provider-specific code. Requirements: 2.12 """ from __future__ import annotations import logging import time import asyncpg from services.recommendation.thesis_llm import ( _log_thesis_performance, _strip_thinking_block, build_thesis_rewrite_prompt, ) from services.shared.agent_config import AgentConfigResolver, ResolvedAgentConfig from services.shared.inference.gateway import InferenceGateway from services.shared.inference.lineage import ModelLineage, build_lineage_from_result from services.shared.inference.models import ( ChatMessage, InferenceTarget, StructuredGenerationRequest, ) from services.shared.schemas import TrendSummary logger = logging.getLogger("recommendation.inference_adapter") async def rewrite_thesis_via_gateway( deterministic_thesis: str, summary: TrendSummary, gateway: InferenceGateway, target: InferenceTarget, pool: asyncpg.Pool | None = None, ) -> tuple[str, ModelLineage]: """Rewrite a deterministic thesis using the InferenceGateway. This replaces the duplicate Ollama/vLLM branching in thesis_llm.py with a single implementation routed through the gateway. If the LLM call fails, returns the original deterministic thesis. The gateway handles protocol selection (Ollama, OpenAI-compatible, etc.) transparently. Args: deterministic_thesis: The rule-based thesis string. summary: The trend summary that produced the thesis. gateway: The shared InferenceGateway instance. target: Resolved inference target for thesis rewriting. pool: Optional asyncpg pool for performance logging. Returns: Tuple of (rewritten thesis, lineage). Falls back to deterministic thesis on failure. """ start_time = time.monotonic() # Resolve agent config for token budget check resolved: ResolvedAgentConfig | None = None if pool is not None: try: resolver = AgentConfigResolver(pool, ttl_seconds=60) resolved = await resolver.resolve("thesis-rewriter") except Exception: logger.warning( "Failed to resolve thesis-rewriter config — proceeding without budget check", exc_info=True, ) # Token budget enforcement if ( resolved is not None and resolved.token_budget > 0 and resolved.variant_id is not None and pool is not None ): try: row = await pool.fetchrow( """SELECT COALESCE(SUM(input_tokens + output_tokens), 0) AS total_tokens FROM agent_performance_log WHERE variant_id = $1 AND recorded_at >= NOW() - INTERVAL '1 hour'""", resolved.variant_id, ) used = int(row["total_tokens"]) if row else 0 if used >= resolved.token_budget: logger.warning( "Token budget exceeded for thesis-rewriter variant %s: used %d / budget %d", resolved.variant_id, used, resolved.token_budget, ) return deterministic_thesis, ModelLineage(model=target.model, protocol=target.protocol) except Exception: logger.warning("Failed to check token budget for thesis-rewriter", exc_info=True) prompts = build_thesis_rewrite_prompt(deterministic_thesis, summary) # Override system prompt from resolved config if resolved is not None and resolved.system_prompt: prompts["system"] = resolved.system_prompt # Build the inference request — no JSON schema needed for thesis rewriting request = StructuredGenerationRequest( messages=[ ChatMessage(role="system", content=prompts["system"]), ChatMessage(role="user", content=prompts["user"]), ], json_schema=None, max_output_tokens=512, temperature=0.0, seed=None, timeout_seconds=target.timeout_seconds, trace_id=f"thesis-{summary.entity_id}", ) try: result = await gateway.generate(target, request) lineage = build_lineage_from_result(result, trace_id=f"thesis-{summary.entity_id}") duration_ms = int((time.monotonic() - start_time) * 1000) if result.error: logger.warning( "LLM thesis rewrite failed for %s: %s — using deterministic thesis", summary.entity_id, result.error, ) if pool is not None and resolved is not None: await _log_thesis_performance( pool, resolved=resolved, ticker=summary.entity_id, success=False, duration_ms=duration_ms, input_tokens=len(deterministic_thesis) // 4, output_tokens=0, error_message=result.error, ) return deterministic_thesis, lineage content = result.content.strip() # Strip thinking blocks (Qwen models) content = _strip_thinking_block(content) if content: logger.info( "LLM thesis rewrite succeeded for %s (%d chars → %d chars)", summary.entity_id, len(deterministic_thesis), len(content), ) if pool is not None and resolved is not None: await _log_thesis_performance( pool, resolved=resolved, ticker=summary.entity_id, success=True, duration_ms=duration_ms, input_tokens=len(deterministic_thesis) // 4, output_tokens=len(content) // 4, ) return content, lineage logger.warning( "LLM thesis rewrite returned empty for %s — using deterministic thesis", summary.entity_id, ) if pool is not None and resolved is not None: await _log_thesis_performance( pool, resolved=resolved, ticker=summary.entity_id, success=False, duration_ms=duration_ms, input_tokens=len(deterministic_thesis) // 4, output_tokens=0, error_message="empty_response", ) return deterministic_thesis, lineage except Exception: duration_ms = int((time.monotonic() - start_time) * 1000) logger.exception( "LLM thesis rewrite failed for %s — using deterministic thesis", summary.entity_id, ) lineage = ModelLineage(model=target.model, protocol=target.protocol) if pool is not None and resolved is not None: await _log_thesis_performance( pool, resolved=resolved, ticker=summary.entity_id, success=False, duration_ms=duration_ms, input_tokens=len(deterministic_thesis) // 4, output_tokens=0, error_message="exception", ) return deterministic_thesis, lineage