1use std::sync::Arc;
2use std::time::Duration;
3
4use agent_base::{ChatMessage, ResponseFormat, StreamClient};
5use serde::de::DeserializeOwned;
6
7pub trait FocusInput {
11 fn to_prompt(&self) -> String;
13}
14
15impl 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
34pub 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 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#[derive(Debug)]
89pub struct FocusOutput<T> {
90 pub result: T,
92 pub raw_response: String,
94}
95
96pub 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 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 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 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 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#[derive(Debug)]
224pub enum FocusError {
225 Timeout(Duration),
227 Llm(String),
229 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#[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 struct MockStreamClient {
263 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 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 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 #[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 #[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 #[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 #[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 #[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 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 assert!(format!("{:?}", focus).contains("Focus"));
475 }
476}