"""Async LLM adapter implementations used by the latency optimisation pipeline.""" from __future__ import annotations import logging from abc import ABC, abstractmethod from typing import TYPE_CHECKING from app.config import settings from app.generator.prompts import ( ENHANCE_SYSTEM_PROMPT, PROOFREAD_SYSTEM_PROMPT, RICS_PROMPT, VALIDATE_SYSTEM_PROMPT, build_enhance_prompt, build_lcel_invoke_vars, build_proofread_prompt, build_validate_prompt, max_context_tokens_for_survey_level, max_output_tokens_for_survey_level, top_p_for_ai_involvement, ) from app.llm.llm_throttle import throttled_llm_call logger = logging.getLogger(__name__) if TYPE_CHECKING: from app.models.schemas import WritingStyleProfile class AsyncLLMAdapter(ABC): @abstractmethod async def generate_section( self, skeleton: str, bullets: list[str], snippets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.2, creativity_hint: str = "", document_context: list[str] | None = None, style_anchor: str | None = None, hierarchy_section_snippets: list[str] | None = None, paragraph_snippets: list[str] | None = None, identity_facts: str | None = None, survey_level: int | None = None, reference_only_context: bool = False, ai_percent: int | None = None, interference_level: str | None = None, scope_fence: str | None = None, ) -> str: ... @abstractmethod async def proofread( self, text: str, bullets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.15, creativity_hint: str = "", ) -> str: ... @abstractmethod async def enhance( self, text: str, bullets: list[str], snippets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.2, creativity_hint: str = "", ) -> str: ... @abstractmethod async def validate_section_compliance( self, *, survey_level: int | None, section_code: str, bullets: list[str], evidence_snippets: list[str], text: str, ) -> str: ... @abstractmethod async def constrained_weave( self, *, section_code: str, section_title: str | None, bullets: list[str], standard_passages: list[str], survey_level: int | None = None, tenant_id: str | None = None, ) -> str: ... class AsyncOpenAIAdapter(AsyncLLMAdapter): """Async generate mode via LangChain LCEL; proofread/enhance/validate via OpenAI ChatCompletions.""" def __init__(self) -> None: from langchain_openai import ChatOpenAI from openai import AsyncOpenAI self._client = AsyncOpenAI(api_key=settings.openai_api_key) self._model = settings.chat_model self._lc_llm = ChatOpenAI( model=self._model, temperature=0.2, max_tokens=settings.max_output_tokens, api_key=settings.openai_api_key, max_retries=3, ) async def _call_async( self, *, system: str, user: str, phase: str, section_id: str | None, max_tokens: int | None = None, temperature: float = 0.2, survey_level: int | None = None, interference_level: str | None = None, tenant_id: str | None = None, ) -> str: from app.llm.llm_throttle import make_cache_hit_slot from app.llm.prompt_cache import ( build_chat_messages, log_openai_cache_usage, openai_extra_kwargs, prompt_caching_active, ) cache_slot = make_cache_hit_slot() async def _do_call() -> str: messages = build_chat_messages(system=system, user=user) extra = openai_extra_kwargs( phase=phase, model=self._model, survey_level=survey_level, interference_level=interference_level, tenant_id=tenant_id, ) response = await self._client.chat.completions.create( model=self._model, messages=messages, max_tokens=max_tokens or settings.max_output_tokens, temperature=temperature, **extra, ) if prompt_caching_active(): cache_slot[0] = log_openai_cache_usage( response, phase=phase, section_id=section_id ) return (response.choices[0].message.content or "").strip() return await throttled_llm_call( phase=phase, section_id=section_id, cache_hit_out=cache_slot, call=_do_call, ) async def generate_section( self, skeleton: str, bullets: list[str], snippets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.2, creativity_hint: str = "", document_context: list[str] | None = None, style_anchor: str | None = None, hierarchy_section_snippets: list[str] | None = None, paragraph_snippets: list[str] | None = None, identity_facts: str | None = None, survey_level: int | None = None, reference_only_context: bool = False, ai_percent: int | None = None, interference_level: str | None = None, tenant_id: str | None = None, scope_fence: str | None = None, ) -> str: from langchain_core.output_parsers import StrOutputParser fine = paragraph_snippets if paragraph_snippets is not None else snippets out_tokens = max_output_tokens_for_survey_level( survey_level, interference_level=interference_level ) ctx_tokens = max_context_tokens_for_survey_level( survey_level, interference_level=interference_level ) vars_ = build_lcel_invoke_vars( skeleton=skeleton, bullets=bullets, snippets=None, max_context_tokens=ctx_tokens, style_profile=style_profile, creativity_hint=creativity_hint, document_snippets=document_context, section_snippets=None, hierarchy_section_snippets=hierarchy_section_snippets, paragraph_snippets=fine, style_anchor=style_anchor, identity_facts=identity_facts, survey_level=survey_level, reference_only_context=reference_only_context, ai_percent=ai_percent, interference_level=interference_level, scope_fence=scope_fence, ) top_p = top_p_for_ai_involvement(ai_percent) chain = ( RICS_PROMPT | self._lc_llm.bind( temperature=temperature, max_tokens=out_tokens, top_p=top_p ) | StrOutputParser() ) from app.llm.prompt_cache import ( build_chat_messages, log_openai_cache_usage, openai_extra_kwargs, prompt_caching_active, ) phase = "generate_section" if prompt_caching_active(): system = str(vars_.get("system_content") or "") user = str(vars_.get("user_content") or "") extra = openai_extra_kwargs( phase=phase, model=self._model, survey_level=survey_level, interference_level=interference_level, tenant_id=tenant_id, ) from app.llm.llm_throttle import make_cache_hit_slot cache_slot = make_cache_hit_slot() async def _cached_generate() -> str: response = await self._client.chat.completions.create( model=self._model, messages=build_chat_messages(system=system, user=user), max_tokens=out_tokens, temperature=temperature, top_p=top_p, **extra, ) cache_slot[0] = log_openai_cache_usage( response, phase=phase, section_id=None ) return (response.choices[0].message.content or "").strip() return await throttled_llm_call( phase=phase, section_id=None, cache_hit_out=cache_slot, call=_cached_generate, ) from app.llm.lcel_invoke import ainvoke_lcel_chain return await ainvoke_lcel_chain( chain, vars_, phase=phase, section_id=None, ) async def proofread( self, text: str, bullets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.15, creativity_hint: str = "", ) -> str: from app.generator.prompts import build_proofread_prompt from app.chunking.splitter import count_tokens user_prompt = build_proofread_prompt( text=text, bullets=bullets, style_profile=style_profile, creativity_hint=creativity_hint, ) out_cap = max(700, min(2400, count_tokens(text or "") + 200)) return await self._call_async( system=PROOFREAD_SYSTEM_PROMPT, user=user_prompt, phase="proofread", section_id=None, max_tokens=out_cap, temperature=temperature, ) async def enhance( self, text: str, bullets: list[str], snippets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.2, creativity_hint: str = "", ) -> str: from app.chunking.splitter import count_tokens user_prompt = build_enhance_prompt( text=text, bullets=bullets, snippets=snippets, max_context_tokens=settings.max_context_tokens, style_profile=style_profile, creativity_hint=creativity_hint, ) out_cap = max(900, min(2400, count_tokens(text or "") + 800)) return await self._call_async( system=ENHANCE_SYSTEM_PROMPT, user=user_prompt, phase="enhance", section_id=None, max_tokens=out_cap, temperature=temperature, ) async def validate_section_compliance( self, *, survey_level: int | None, section_code: str, bullets: list[str], evidence_snippets: list[str], text: str, ) -> str: user_prompt = build_validate_prompt( survey_level=survey_level, section_code=section_code, bullets=bullets, evidence_snippets=evidence_snippets, text=text, ) result = await self._call_async( system=VALIDATE_SYSTEM_PROMPT, user=user_prompt, phase="validate_section", section_id=section_code, max_tokens=220, temperature=0.0, ) return (result or "").strip() async def constrained_weave( self, *, section_code: str, section_title: str | None, bullets: list[str], standard_passages: list[str], survey_level: int | None = None, tenant_id: str | None = None, ) -> str: from app.generator.prompts import _ASSEMBLY_SYSTEM_CORE cleaned_passages = [ str(p).strip() for p in (standard_passages or []) if str(p).strip() ] cleaned_bullets = [str(b).strip() for b in (bullets or []) if str(b).strip()] if not cleaned_passages or not cleaned_bullets: return "" title_part = f" — {section_title}" if section_title else "" user = ( f"SECTION: {section_code}{title_part}\n\n" "STANDARD SOURCE PASSAGES (preserve wording; weave NOTES facts into the slots):\n" + "\n".join(f"- {p}" for p in cleaned_passages) + "\n\nINSPECTOR'S RAW NOTES (substitute these specifics into the standards):\n" + "\n".join(f"- {b}" for b in cleaned_bullets) + "\n\nProduce the structurally-routed output now. Standard wording stays, " "note facts replace generic slots, no new sentences, no new claims." ) from app.llm.llm_throttle import make_cache_hit_slot from app.llm.prompt_cache import ( build_chat_messages, log_openai_cache_usage, openai_extra_kwargs, prompt_caching_active, ) cache_slot = make_cache_hit_slot() phase = "constrained_weave" async def _do_call() -> str: msgs = build_chat_messages(system=_ASSEMBLY_SYSTEM_CORE, user=user) extra = openai_extra_kwargs( phase=phase, model=self._model, survey_level=survey_level, tenant_id=tenant_id, ) response = await self._client.chat.completions.create( model=self._model, messages=msgs, max_tokens=600, temperature=0.0, top_p=0.1, **extra, ) if prompt_caching_active(): cache_slot[0] = log_openai_cache_usage( response, phase=phase, section_id=section_code ) return (response.choices[0].message.content or "").strip() return await throttled_llm_call( phase=phase, section_id=section_code, cache_hit_out=cache_slot, call=_do_call, ) class MockAsyncLLMAdapter(AsyncLLMAdapter): """Deterministic mock adapter for async paths without OpenAI.""" def __init__(self) -> None: pass async def generate_section( self, skeleton: str, bullets: list[str], snippets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.2, creativity_hint: str = "", document_context: list[str] | None = None, style_anchor: str | None = None, hierarchy_section_snippets: list[str] | None = None, paragraph_snippets: list[str] | None = None, identity_facts: str | None = None, survey_level: int | None = None, reference_only_context: bool = False, ai_percent: int | None = None, interference_level: str | None = None, scope_fence: str | None = None, ) -> str: style_note = f" (style: {style_profile.tone})" if style_profile else "" summary = "; ".join(bullets[:3]) if bullets else "No facts provided" return f"Based on the available information{style_note}: {summary}." async def proofread( self, text: str, bullets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.15, creativity_hint: str = "", ) -> str: return ( f"{text}\n---NOTES---\n" "No OpenAI key configured — proofreading not available in mock mode." ) async def enhance( self, text: str, bullets: list[str], snippets: list[str], style_profile: "WritingStyleProfile | None" = None, temperature: float = 0.2, creativity_hint: str = "", ) -> str: extra = ( f" Additional context from {len(snippets)} retrieved source(s) noted." if snippets else "" ) return ( f"{text}{extra} " "[No OpenAI key configured — full technical enhancement requires OPENAI_API_KEY.]" ) async def validate_section_compliance( self, *, survey_level: int | None, section_code: str, bullets: list[str], evidence_snippets: list[str], text: str, ) -> str: return "PASS" async def constrained_weave( self, *, section_code: str, section_title: str | None, bullets: list[str], standard_passages: list[str], survey_level: int | None = None, tenant_id: str | None = None, ) -> str: return "" _async_llm_adapter_instance: AsyncLLMAdapter | None = None def get_async_llm_adapter() -> AsyncLLMAdapter: """Return a singleton async adapter (real OpenAI when key configured).""" global _async_llm_adapter_instance if _async_llm_adapter_instance is not None: return _async_llm_adapter_instance if settings.openai_api_key: _async_llm_adapter_instance = AsyncOpenAIAdapter() else: _async_llm_adapter_instance = MockAsyncLLMAdapter() return _async_llm_adapter_instance