Skip to main content

agent_works/focus/
core.rs

1use std::sync::Arc;
2use std::time::Duration;
3
4use agent_base::{ChatMessage, ResponseFormat, StreamClient};
5use serde::de::DeserializeOwned;
6
7// ── FocusInput ───────────────────────────────────────────────────────────────
8
9/// Input for a Focus call — either a simple string or a structured context.
10pub trait FocusInput {
11    /// Format the input into the user prompt text sent to the LLM.
12    fn to_prompt(&self) -> String;
13}
14
15/// Simple case: pass a string directly.
16impl FocusInput for str {
17    fn to_prompt(&self) -> String {
18        self.to_string()
19    }
20}
21
22impl FocusInput for String {
23    fn to_prompt(&self) -> String {
24        self.clone()
25    }
26}
27
28impl FocusInput for &str {
29    fn to_prompt(&self) -> String {
30        self.to_string()
31    }
32}
33
34// ── Context ──────────────────────────────────────────────────────────────────
35
36/// Structured context for multi-field input scenarios.
37///
38/// Fields are formatted as `【key】\nvalue` when sent to the LLM,
39/// where the key acts as a label to help the LLM understand the context.
40///
41/// # Usage
42///
43/// ```ignore
44/// let ctx = Context::new()
45///     .add("command", "apt install nginx")
46///     .add("screen", screen_content);
47/// ```
48pub struct Context {
49    entries: Vec<(String, String)>,
50}
51
52impl Context {
53    pub fn new() -> Self {
54        Self {
55            entries: Vec::new(),
56        }
57    }
58
59    /// Add a context field. The key is used as a label when sent to the LLM.
60    pub fn add(mut self, key: &str, value: &str) -> Self {
61        self.entries.push((key.to_string(), value.to_string()));
62        self
63    }
64}
65
66impl Default for Context {
67    fn default() -> Self {
68        Self::new()
69    }
70}
71
72impl FocusInput for Context {
73    fn to_prompt(&self) -> String {
74        self.entries
75            .iter()
76            .map(|(key, value)| format!("【{}】\n{}", key, value))
77            .collect::<Vec<_>>()
78            .join("\n\n")
79    }
80}
81
82// ── FocusOutput ──────────────────────────────────────────────────────────────
83
84/// Output wrapper for a Focus call.
85///
86/// Contains both the structured result and the raw LLM response,
87/// useful for debugging when something goes wrong.
88#[derive(Debug)]
89pub struct FocusOutput<T> {
90    /// Deserialized structured result.
91    pub result: T,
92    /// Raw LLM response text (JSON string), for logging and debugging.
93    pub raw_response: String,
94}
95
96// ── Focus ────────────────────────────────────────────────────────────────────
97
98/// A focused LLM call.
99///
100/// Each instance is bound to a system prompt and dedicated to one specific
101/// judgment question. Use `ask()` to send input and receive a structured
102/// JSON answer.
103///
104/// # Usage
105///
106/// ```ignore
107/// // Simple case: single string input
108/// let classify = Focus::new(client, "You are a task complexity classifier...");
109/// let output = classify.ask::<TaskComplexity>(&user_input, 5s).await?;
110///
111/// // Complex case: multiple context fields
112/// let status_focus = Focus::new(client, "You are a task status judge...");
113/// let ctx = Context::new()
114///     .add("command", command)
115///     .add("screen", screen);
116/// let output = status_focus.ask::<TaskStatus>(&ctx, 5s).await?;
117/// ```
118pub struct Focus {
119    client: Arc<dyn StreamClient>,
120    system_prompt: String,
121}
122
123impl std::fmt::Debug for Focus {
124    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
125        f.debug_struct("Focus").finish_non_exhaustive()
126    }
127}
128
129impl Focus {
130    /// Create a new Focus instance.
131    ///
132    /// - `client`: LLM client (shared; multiple Focus instances can reuse the same client)
133    /// - `system_prompt`: The role and judgment rules for this Focus (bound at creation, never changes)
134    pub fn new(client: Arc<dyn StreamClient>, system_prompt: impl Into<String>) -> Self {
135        Self {
136            client,
137            system_prompt: system_prompt.into(),
138        }
139    }
140
141    /// Make a focused LLM call.
142    ///
143    /// Sends the system prompt (bound at creation) + user input (this call),
144    /// forces JSON output, and deserializes into `T`.
145    ///
146    /// # Arguments
147    /// - `input`: User input — can be `&str` or `Context`
148    /// - `timeout`: Call timeout
149    ///
150    /// # Returns
151    /// `FocusOutput<T>` containing the structured result and raw response.
152    pub async fn ask<T: DeserializeOwned>(
153        &self,
154        input: &(impl FocusInput + ?Sized),
155        timeout: Duration,
156    ) -> Result<FocusOutput<T>, FocusError> {
157        let user_prompt = input.to_prompt();
158
159        // Logging: first line of prompt + char count (privacy-friendly)
160        let prompt_first_line = user_prompt.lines().next().unwrap_or("(empty)");
161        let prompt_char_count = user_prompt.chars().count();
162        let sys_first_line = self.system_prompt.lines().next().unwrap_or("(empty)");
163        let target_type = std::any::type_name::<T>();
164
165        tracing::info!(
166            target_type = target_type,
167            system_prompt = %sys_first_line,
168            user_prompt_first_line = %prompt_first_line,
169            user_prompt_chars = prompt_char_count,
170            timeout_secs = timeout.as_secs(),
171            "[Focus] calling LLM"
172        );
173
174        let start = std::time::Instant::now();
175        let messages = vec![
176            ChatMessage::system(self.system_prompt.clone()),
177            ChatMessage::user(user_prompt),
178        ];
179
180        let response = tokio::time::timeout(
181            timeout,
182            self.client
183                .chat(&messages, &[], None, Some(&ResponseFormat::JsonObject)),
184        )
185        .await
186        .map_err(|_| FocusError::Timeout(timeout))?
187        .map_err(|e| FocusError::Llm(e.to_string()))?;
188
189        let elapsed_ms = start.elapsed().as_millis();
190        // StreamClient::chat() returns extracted text — no JSON unwrapping needed.
191        let raw_response = response;
192
193        let result: T = serde_json::from_str(&raw_response).map_err(|e| {
194            tracing::warn!(
195                error = %e,
196                raw_response = %raw_response,
197                elapsed_ms = elapsed_ms,
198                "[Focus] failed to parse LLM response as JSON"
199            );
200            FocusError::Parse {
201                error: e.to_string(),
202                raw: raw_response.clone(),
203            }
204        })?;
205
206        tracing::info!(
207            target_type = target_type,
208            raw_response_chars = raw_response.chars().count(),
209            elapsed_ms = elapsed_ms,
210            "[Focus] call succeeded"
211        );
212
213        Ok(FocusOutput {
214            result,
215            raw_response,
216        })
217    }
218}
219
220// ── FocusError ───────────────────────────────────────────────────────────────
221
222/// Error type for Focus calls.
223#[derive(Debug)]
224pub enum FocusError {
225    /// LLM call timed out.
226    Timeout(Duration),
227    /// LLM call failed (network error, API error, etc.).
228    Llm(String),
229    /// LLM response could not be parsed into the expected JSON type.
230    Parse { error: String, raw: String },
231}
232
233impl std::fmt::Display for FocusError {
234    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
235        match self {
236            FocusError::Timeout(d) => write!(f, "Focus timeout after {:?}", d),
237            FocusError::Llm(e) => write!(f, "Focus LLM error: {}", e),
238            FocusError::Parse { error, .. } => write!(f, "Focus parse error: {}", error),
239        }
240    }
241}
242
243impl std::error::Error for FocusError {}
244
245// ── Tests ────────────────────────────────────────────────────────────────────
246
247#[cfg(test)]
248mod tests {
249    use super::*;
250    use agent_base::{LlmCapabilities, StreamChunk};
251    use async_trait::async_trait;
252    use futures_core::Stream;
253    use serde::Deserialize;
254    use std::pin::Pin;
255    use std::sync::Mutex;
256    use std::task::Context as TaskContext;
257    use std::task::Poll;
258
259    // ── Mock StreamClient for Focus tests ──
260
261    /// A mock StreamClient whose `chat()` returns a pre-set string.
262    struct MockStreamClient {
263        /// Canned response for `chat()`. Consumed on first call (take).
264        response: Mutex<Option<Result<String, String>>>,
265    }
266
267    impl MockStreamClient {
268        fn with_text(text: impl Into<String>) -> Self {
269            Self {
270                response: Mutex::new(Some(Ok(text.into()))),
271            }
272        }
273
274        fn with_error(err: impl Into<String>) -> Self {
275            Self {
276                response: Mutex::new(Some(Err(err.into()))),
277            }
278        }
279    }
280
281    /// Empty stream — returned by MockStreamClient::stream() which is never called.
282    struct EmptyStream;
283
284    impl Stream for EmptyStream {
285        type Item = agent_base::AgentResult<StreamChunk>;
286        fn poll_next(self: Pin<&mut Self>, _cx: &mut TaskContext<'_>) -> Poll<Option<Self::Item>> {
287            Poll::Ready(None)
288        }
289    }
290
291    #[async_trait]
292    impl StreamClient for MockStreamClient {
293        async fn stream(
294            &self,
295            _messages: &[ChatMessage],
296            _tools: &[serde_json::Value],
297            _reasoning: Option<&agent_base::ReasoningConfig>,
298            _response_format: Option<&ResponseFormat>,
299        ) -> agent_base::AgentResult<
300            Pin<Box<dyn Stream<Item = agent_base::AgentResult<StreamChunk>> + Send>>,
301        > {
302            Ok(Box::pin(EmptyStream))
303        }
304
305        /// Override `chat()` to return the canned response directly.
306        async fn chat(
307            &self,
308            _messages: &[ChatMessage],
309            _tools: &[serde_json::Value],
310            _reasoning: Option<&agent_base::ReasoningConfig>,
311            _response_format: Option<&ResponseFormat>,
312        ) -> agent_base::AgentResult<String> {
313            match self.response.lock().unwrap().take() {
314                Some(Ok(text)) => Ok(text),
315                Some(Err(e)) => Err(agent_base::AgentError::internal(e)),
316                None => Ok(String::new()),
317            }
318        }
319
320        fn capabilities(&self) -> LlmCapabilities {
321            LlmCapabilities::default()
322        }
323    }
324
325    // ── Context tests ──
326
327    #[test]
328    fn context_single_field() {
329        let ctx = Context::new().add("command", "df -h");
330        assert_eq!(ctx.to_prompt(), "【command】\ndf -h");
331    }
332
333    #[test]
334    fn context_multiple_fields() {
335        let ctx = Context::new()
336            .add("command", "apt install nginx")
337            .add("elapsed", "30s")
338            .add("screen", "Reading package lists...");
339        let expected = "【command】\napt install nginx\n\n【elapsed】\n30s\n\n【screen】\nReading package lists...";
340        assert_eq!(ctx.to_prompt(), expected);
341    }
342
343    #[test]
344    fn context_empty() {
345        let ctx = Context::new();
346        assert_eq!(ctx.to_prompt(), "");
347    }
348
349    // ── FocusInput tests ──
350
351    #[test]
352    fn str_input() {
353        let input: &str = "hello";
354        assert_eq!(input.to_prompt(), "hello");
355    }
356
357    #[test]
358    fn string_input() {
359        let input = String::from("hello");
360        assert_eq!(input.to_prompt(), "hello");
361    }
362
363    // ── FocusOutput deserialize tests ──
364
365    #[derive(Deserialize, Debug, PartialEq)]
366    struct MockResult {
367        status: String,
368        reason: String,
369    }
370
371    #[test]
372    fn focus_output_deserialize() {
373        let raw = r#"{"status":"finished","reason":"done"}"#;
374        let result: MockResult = serde_json::from_str(raw).unwrap();
375        assert_eq!(result.status, "finished");
376        assert_eq!(result.reason, "done");
377    }
378
379    // ── FocusError tests ──
380
381    #[test]
382    fn focus_error_display() {
383        let err = FocusError::Timeout(Duration::from_secs(5));
384        assert_eq!(format!("{}", err), "Focus timeout after 5s");
385
386        let err = FocusError::Llm("network error".to_string());
387        assert_eq!(format!("{}", err), "Focus LLM error: network error");
388
389        let err = FocusError::Parse {
390            error: "unexpected token".to_string(),
391            raw: "not json".to_string(),
392        };
393        assert_eq!(format!("{}", err), "Focus parse error: unexpected token");
394    }
395
396    // ── Focus::ask() tests ──
397
398    #[derive(Deserialize, Debug, PartialEq)]
399    struct AskResult {
400        status: String,
401        confidence: f64,
402    }
403
404    #[tokio::test]
405    async fn focus_ask_parses_valid_json() {
406        let client = Arc::new(MockStreamClient::with_text(
407            r#"{"status":"finished","confidence":0.95}"#,
408        ));
409        let focus = Focus::new(client, "You are a classifier.");
410        let output: FocusOutput<AskResult> = focus
411            .ask(&"classify this", Duration::from_secs(5))
412            .await
413            .expect("ask should succeed");
414        assert_eq!(output.result.status, "finished");
415        assert_eq!(output.result.confidence, 0.95);
416        assert_eq!(
417            output.raw_response,
418            r#"{"status":"finished","confidence":0.95}"#
419        );
420    }
421
422    #[tokio::test]
423    async fn focus_ask_parses_str_input() {
424        let client = Arc::new(MockStreamClient::with_text(
425            r#"{"status":"done","confidence":1.0}"#,
426        ));
427        let focus = Focus::new(client, "system");
428        let output: FocusOutput<AskResult> = focus
429            .ask("classify", Duration::from_secs(5))
430            .await
431            .expect("ask should succeed");
432        assert_eq!(output.result.status, "done");
433    }
434
435    #[tokio::test]
436    async fn focus_ask_rejects_invalid_json() {
437        let client = Arc::new(MockStreamClient::with_text("not valid json at all"));
438        let focus = Focus::new(client, "system");
439        let err = focus
440            .ask::<AskResult>(&"input", Duration::from_secs(5))
441            .await
442            .unwrap_err();
443        assert!(matches!(err, FocusError::Parse { .. }));
444        assert!(err.to_string().contains("Focus parse error"));
445    }
446
447    #[tokio::test]
448    async fn focus_ask_propagates_llm_error() {
449        let client = Arc::new(MockStreamClient::with_error("api key invalid"));
450        let focus = Focus::new(client, "system");
451        let err = focus
452            .ask::<AskResult>(&"input", Duration::from_secs(5))
453            .await
454            .unwrap_err();
455        assert!(matches!(err, FocusError::Llm(_)));
456    }
457
458    #[tokio::test]
459    async fn focus_ask_times_out() {
460        // Return after a long delay — Focus has a very short timeout
461        let client = Arc::new(MockStreamClient::with_text("{}"));
462        let focus = Focus::new(client, "system");
463        let result = focus
464            .ask::<AskResult>(&"input", Duration::from_millis(1))
465            .await;
466        assert!(result.is_err());
467    }
468
469    #[tokio::test]
470    async fn focus_ask_new_constructs_correctly() {
471        let client = Arc::new(MockStreamClient::with_text("{}"));
472        let focus = Focus::new(client, "You are helpful.");
473        // Just verify construction + Debug
474        assert!(format!("{:?}", focus).contains("Focus"));
475    }
476}