1"""DraftVerifyExecutor — draft-then-verify pattern using cheap and expensive models."""
2
3from __future__ import annotations
4
5import asyncio
6
7from lexigram.contracts import (
8 LLMClientProtocol,
9)
10from lexigram.contracts.ai.exceptions import LLMError
11from lexigram.contracts.ai.llm import ChatMessage, Completion
12from lexigram.logging import (
13 get_logger,
14)
15from lexigram.result import Ok, Result
16
17logger = get_logger(__name__)
18
19
20class DraftVerifyExecutor:
21 """Draft-then-verify pattern using cheap + expensive model pair.
22
23 Fires a cheap/fast model and an expensive/capable model in parallel.
24 If the cheap model's draft passes verification, cancels the expensive
25 model and returns the draft immediately. Otherwise waits for the
26 expensive model's result.
27
28 Opt-in strategy for scenarios where latency matters more than guaranteed
29 quality on the first try.
30 """
31
32 def __init__(
33 self,
34 draft_client: LLMClientProtocol,
35 verify_client: LLMClientProtocol,
36 pro_client: LLMClientProtocol,
37 ) -> None:
38 """Initialize the DraftVerifyExecutor.
39
40 Args:
41 draft_client: Fast/cheap LLM client for draft generation.
42 verify_client: Verifier LLM client (typically small/fast).
43 pro_client: Slow/expensive LLM client for fallback.
44 """
45 self._draft = draft_client
46 self._verify = verify_client
47 self._pro = pro_client
48 self._background_tasks: set[asyncio.Task] = set()
49
50 async def execute(
51 self,
52 messages: list[ChatMessage],
53 model: str | None = None,
54 temperature: float | None = None,
55 max_tokens: int | None = None,
56 ) -> Result[Completion, LLMError]:
57 """Execute draft-then-verify pattern.
58
59 1. Fire draft_client and pro_client in parallel.
60 2. Await draft (faster).
61 3. Verify draft with verify_client (yes/no).
62 4. If verified: cancel pro_client, return draft.
63 5. If not verified: await pro_client, return its result.
64
65 Args:
66 messages: Chat messages to send to both models.
67 model: Optional model override (applied to pro_client only).
68 temperature: Optional temperature override.
69 max_tokens: Optional max_tokens override.
70
71 Returns:
72 Result containing Completion on success, LLMError on failure.
73 """
74 draft_task = asyncio.create_task(
75 self._draft.complete(
76 messages,
77 temperature=temperature,
78 max_tokens=max_tokens,
79 )
80 )
81 pro_task = asyncio.create_task(
82 self._pro.complete(
83 messages,
84 model=model,
85 temperature=temperature,
86 max_tokens=max_tokens,
87 )
88 )
89 self._background_tasks.add(draft_task)
90 self._background_tasks.add(pro_task)
91 draft_task.add_done_callback(self._background_tasks.discard)
92 pro_task.add_done_callback(self._background_tasks.discard)
93
94 draft_result = await draft_task
95 if draft_result.is_ok():
96 draft_completion = draft_result.unwrap()
97 else:
98 pro_task.cancel()
99 return draft_result # type: ignore[return-value]
100 draft_text = draft_completion.content
101
102 verify_messages = [
103 *messages,
104 ChatMessage(role="assistant", content=draft_text),
105 ChatMessage(
106 role="user",
107 content=f"Is the following response correct and complete? Response: {draft_text[:200]}...",
108 ),
109 ]
110 verify_result = await self._verify.complete(verify_messages, max_tokens=10)
111 if verify_result.is_ok():
112 verify_text = verify_result.unwrap().content.lower()
113 else:
114 logger.warning("draft_verify_failed", error=str(verify_result.unwrap_err()))
115 verify_text = "invalid"
116
117 rejected = any(
118 w in verify_text for w in ["no", "not", "incorrect", "invalid", "wrong"]
119 )
120 accepted = any(w in verify_text for w in ["yes", "correct", "valid"])
121 draft_accepted = accepted and not rejected
122
123 if draft_accepted:
124 pro_task.cancel()
125 logger.info("draft_verify_accepted")
126 return Ok(draft_completion) # type: ignore[arg-type]
127
128 logger.info("draft_verify_rejected")
129 return await pro_task # type: ignore[return-value]