kcode_k1_codex_websearch/
lib.rs1use 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}