"""LLM parser using Ollama: cache lookup → prompt → validate → store.""" from __future__ import annotations import json from typing import Any from sqlalchemy.ext.asyncio import AsyncSession from libs.common.logging import get_logger from libs.llm.cache import LLMCacheStore, build_cache_key from libs.llm.client import OllamaClient from libs.llm.exceptions import LLMError, LLMSchemaError from libs.llm.prompts import PROMPT_VERSION, render_prompt from libs.parser.schema_validator import validate_parser_output logger = get_logger(__name__) _SCHEMA_VERSION = "1.0.0" _PROMPT_NAME = "event_classifier_v1" # Fields returned by LLM that map to ParserEventOutput fields _REQUIRED_FIELDS = { "event_type", "event_direction", "headline", "summary", "guidance_status", "confidence_overall", } def _llm_response_to_parser_output( document_id: str, doc_meta: dict[str, Any], raw: dict[str, Any], ) -> dict[str, Any]: """Convert flat LLM response dict → ParserEventOutput-compatible dict.""" return { "schema_version": _SCHEMA_VERSION, "document_id": document_id, "parser_kind": "llm", "event_type": raw.get("event_type", "unknown"), "event_direction": raw.get("event_direction", "unknown"), "event_date": doc_meta.get("filing_date", "1900-01-01"), "filing_time_bucket": doc_meta.get("filing_time_bucket", "unknown"), "headline": raw.get("headline", ""), "summary": raw.get("summary", ""), "guidance": { "status": raw.get("guidance_status", "unclear"), "scope": raw.get("guidance_scope", "unknown"), "notes": "", }, "signals": { "demand_strength": raw.get("demand_strength", "unknown"), "pricing_power": raw.get("pricing_power", "unknown"), "backlog_or_bookings": raw.get("backlog_or_bookings", "unknown"), "customer_expansion": raw.get("customer_expansion", "unknown"), "margin_quality": raw.get("margin_quality", "unknown"), }, "risk_flags": { "oneoff_item": bool(raw.get("oneoff_item", False)), "tax_benefit": bool(raw.get("tax_benefit", False)), "valuation_gain": bool(raw.get("valuation_gain", False)), "non_gaap_heavy": bool(raw.get("non_gaap_heavy", False)), "financing_related": bool(raw.get("financing_related", False)), "legal_or_regulatory_overhang": bool( raw.get("legal_or_regulatory_overhang", False) ), }, "evidence": [], "confidence": { "overall": float(raw.get("confidence_overall", 0.5)), "event_type": float(raw.get("confidence_event_type", 0.5)), "event_direction": float(raw.get("confidence_event_direction", 0.5)), "guidance": 0.5, "risk_flags": 0.5, }, "warnings": [], } def _repair_prompt( original_messages: list[dict[str, str]], errors: list[str], raw_response: str, ) -> list[dict[str, str]]: """Append a repair instruction to messages asking LLM to fix schema errors.""" repair_msg = ( f"Your previous response had schema errors: {errors[:3]}. " f"Previous response was: {raw_response[:300]}. " "Please fix these issues and return only valid JSON matching the schema." ) return [*original_messages, {"role": "assistant", "content": raw_response}, {"role": "user", "content": repair_msg}] class LLMParser: """Parse a document with Ollama, using DB cache to avoid redundant calls.""" def __init__( self, client: OllamaClient, cache_store: LLMCacheStore | None = None, ) -> None: self._client = client self._cache = cache_store or LLMCacheStore() async def parse( self, document_id: str, doc_text: str, doc_meta: dict[str, Any], rule_hints: dict[str, Any], session: AsyncSession, ) -> dict[str, Any] | None: """Parse a document with LLM. Returns a ParserEventOutput-compatible dict or None on failure. """ cache_key = build_cache_key( text=doc_text, prompt_name=_PROMPT_NAME, prompt_version=PROMPT_VERSION, model_name=self._client.model, schema_version=_SCHEMA_VERSION, ) # 1. Cache lookup cached = await self._cache.get(session, cache_key) if cached is not None: logger.info("llm_parse_cache_hit", document_id=document_id) return cached # 2. Render prompt messages = render_prompt( prompt_name=_PROMPT_NAME, doc_meta=doc_meta, text=doc_text, rule_hints=rule_hints, ) raw_prompt_str = json.dumps(messages) # 3. Call Ollama try: raw_dict, token_usage, elapsed_ms = await self._client.chat(messages) except LLMError as exc: logger.error("llm_parse_failed", document_id=document_id, error=str(exc)) return None raw_response_str = json.dumps(raw_dict) # 4. Map to parser output format output = _llm_response_to_parser_output(document_id, doc_meta, raw_dict) # 5. Validate schema errors = validate_parser_output(output) if errors: logger.warning("llm_schema_errors_attempt_repair", errors=errors[:3]) repair_messages = _repair_prompt(messages, errors, raw_response_str) try: raw_dict2, token_usage2, elapsed_ms2 = await self._client.chat(repair_messages) token_usage = { "prompt_tokens": token_usage.get("prompt_tokens", 0) + token_usage2.get("prompt_tokens", 0), "completion_tokens": token_usage.get("completion_tokens", 0) + token_usage2.get("completion_tokens", 0), } elapsed_ms += elapsed_ms2 raw_response_str = json.dumps(raw_dict2) output = _llm_response_to_parser_output(document_id, doc_meta, raw_dict2) errors = validate_parser_output(output) except LLMError as exc: logger.error("llm_repair_failed", document_id=document_id, error=str(exc)) return None if errors: logger.error( "llm_schema_invalid_after_repair", document_id=document_id, errors=errors[:3], ) raise LLMSchemaError( f"LLM output invalid after repair: {errors[:2]}", source="llm_parser", context={"document_id": document_id}, ) # 6. Store in cache try: await self._cache.put( session=session, cache_key=cache_key, document_id=document_id, model_name=self._client.model, prompt_version=PROMPT_VERSION, schema_version=_SCHEMA_VERSION, raw_prompt=raw_prompt_str, raw_response=raw_response_str, normalized=output, token_usage=token_usage, elapsed_ms=elapsed_ms, ) except Exception as exc: logger.warning("llm_cache_write_failed", error=str(exc)) logger.info( "llm_parse_ok", document_id=document_id, elapsed_ms=elapsed_ms, event_type=output.get("event_type"), ) return output