Skip to main content

kcode_k1_codex_websearch/
lib.rs

1use kcode_k1_codex_adapter::{Adapter, Error, ErrorKind, Event, SearchTurn};
2use std::{
3    sync::{
4        Arc,
5        atomic::{AtomicU64, Ordering},
6    },
7    time::Instant,
8};
9use tokio::time::{Instant as TokioInstant, timeout_at};
10
11const DEADLINE_ERROR: &str = "WebSearch failed: deadline exceeded";
12const EXHAUSTED_ERROR: &str = "WebSearch failed: operation counter exhausted";
13const EMPTY_ERROR: &str = "WebSearch failed: empty assistant response";
14const TOOL_ERROR: &str = "WebSearch failed: unexpected tool call";
15const STREAM_ERROR: &str = "WebSearch failed: response stream closed";
16
17pub struct Request {
18    pub query: String,
19    pub model: String,
20    pub reasoning_effort: String,
21    pub deadline: Instant,
22}
23
24#[derive(Clone)]
25pub struct Runner {
26    adapter: Adapter,
27    counter: Arc<AtomicU64>,
28}
29
30impl Runner {
31    pub fn new(adapter: Adapter) -> Self {
32        Self {
33            adapter,
34            counter: Arc::new(AtomicU64::new(0)),
35        }
36    }
37
38    pub async fn run(&self, request: Request) -> Result<String, String> {
39        let Request {
40            query,
41            model,
42            reasoning_effort,
43            deadline,
44        } = request;
45        if Instant::now() >= deadline {
46            return Err(DEADLINE_ERROR.into());
47        }
48
49        let key = next_key(&self.counter)?;
50        let start = self
51            .adapter
52            .start_web_search_turn(key, query, model, reasoning_effort);
53        let mut turn = match timeout_at(TokioInstant::from_std(deadline), start).await {
54            Ok(Ok(turn)) => turn,
55            Ok(Err(error)) => return Err(safe_runtime_error(&error)),
56            Err(_) => return Err(DEADLINE_ERROR.into()),
57        };
58
59        let mut text = String::new();
60        loop {
61            let event = match timeout_at(TokioInstant::from_std(deadline), turn.next_event()).await
62            {
63                Ok(event) => event,
64                Err(_) => return Err(DEADLINE_ERROR.into()),
65            };
66            match event {
67                Some(Event::TextDelta(delta)) => text.push_str(&delta),
68                Some(Event::Done) => {
69                    if let Err(error) = require_message(&text) {
70                        return fail_after_close(turn, deadline, error).await;
71                    }
72                    close_turn(turn, deadline).await?;
73                    return Ok(text);
74                }
75                Some(Event::Error(error)) => {
76                    let error = safe_runtime_error(&error);
77                    return fail_after_close(turn, deadline, error).await;
78                }
79                Some(Event::ToolCall(_)) => {
80                    return fail_after_close(turn, deadline, TOOL_ERROR.into()).await;
81                }
82                None => {
83                    return fail_after_close(turn, deadline, STREAM_ERROR.into()).await;
84                }
85            }
86        }
87    }
88}
89
90fn next_key(counter: &AtomicU64) -> Result<String, String> {
91    let previous = counter
92        .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |value| {
93            value.checked_add(1)
94        })
95        .map_err(|_| EXHAUSTED_ERROR.to_owned())?;
96    let number = previous
97        .checked_add(1)
98        .ok_or_else(|| EXHAUSTED_ERROR.to_owned())?;
99    Ok(format!("k1-websearch:{number:016}"))
100}
101
102fn require_message(text: &str) -> Result<(), String> {
103    (!text.is_empty())
104        .then_some(())
105        .ok_or_else(|| EMPTY_ERROR.to_owned())
106}
107
108async fn close_turn(turn: SearchTurn, deadline: Instant) -> Result<(), String> {
109    match timeout_at(TokioInstant::from_std(deadline), turn.close()).await {
110        Ok(Ok(())) => Ok(()),
111        Ok(Err(error)) => Err(safe_runtime_error(&error)),
112        Err(_) => Err(DEADLINE_ERROR.into()),
113    }
114}
115
116async fn fail_after_close(
117    turn: SearchTurn,
118    deadline: Instant,
119    error: String,
120) -> Result<String, String> {
121    match close_turn(turn, deadline).await {
122        Ok(()) => Err(error),
123        Err(close_error) => Err(close_error),
124    }
125}
126
127fn safe_runtime_error(error: &Error) -> String {
128    match error.kind {
129        ErrorKind::Busy => "WebSearch failed: runtime busy",
130        ErrorKind::Interrupted => "WebSearch failed: interrupted",
131        ErrorKind::InvalidToolResult | ErrorKind::Protocol => {
132            "WebSearch failed: runtime protocol error"
133        }
134        ErrorKind::LaunchRejected => "WebSearch failed: launch rejected",
135        ErrorKind::Server => "WebSearch failed: provider error",
136        ErrorKind::Unavailable => "WebSearch failed: runtime unavailable",
137    }
138    .into()
139}
140
141#[cfg(test)]
142mod tests {
143    use super::*;
144
145    #[test]
146    fn shared_counter_produces_exact_unique_keys_and_stops_at_exhaustion() {
147        let counter = Arc::new(AtomicU64::new(0));
148        let clone = Arc::clone(&counter);
149        assert_eq!(next_key(&counter).unwrap(), "k1-websearch:0000000000000001");
150        assert_eq!(next_key(&clone).unwrap(), "k1-websearch:0000000000000002");
151
152        counter.store(u64::MAX, Ordering::Relaxed);
153        assert_eq!(next_key(&counter).unwrap_err(), EXHAUSTED_ERROR);
154        assert_eq!(counter.load(Ordering::Relaxed), u64::MAX);
155    }
156
157    #[test]
158    fn runtime_errors_map_to_safe_stable_categories() {
159        for (kind, expected) in [
160            (ErrorKind::Busy, "WebSearch failed: runtime busy"),
161            (ErrorKind::Interrupted, "WebSearch failed: interrupted"),
162            (
163                ErrorKind::InvalidToolResult,
164                "WebSearch failed: runtime protocol error",
165            ),
166            (
167                ErrorKind::LaunchRejected,
168                "WebSearch failed: launch rejected",
169            ),
170            (
171                ErrorKind::Protocol,
172                "WebSearch failed: runtime protocol error",
173            ),
174            (ErrorKind::Server, "WebSearch failed: provider error"),
175            (
176                ErrorKind::Unavailable,
177                "WebSearch failed: runtime unavailable",
178            ),
179        ] {
180            let error = Error::new(kind, "private detail").with_diagnostics(vec![1, 2, 3]);
181            assert_eq!(safe_runtime_error(&error), expected);
182        }
183    }
184
185    #[test]
186    fn completed_response_requires_only_nonempty_verbatim_text() {
187        assert_eq!(require_message("").unwrap_err(), EMPTY_ERROR);
188        assert!(require_message(" ").is_ok());
189        assert!(require_message("answer\nhttps://example.com").is_ok());
190    }
191}