use kcode_k1_codex_adapter::{Adapter, Error, ErrorKind, Event, SearchTurn};
use std::{
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
time::Instant,
};
use tokio::time::{Instant as TokioInstant, timeout_at};
const DEADLINE_ERROR: &str = "WebSearch failed: deadline exceeded";
const EXHAUSTED_ERROR: &str = "WebSearch failed: operation counter exhausted";
const EMPTY_ERROR: &str = "WebSearch failed: empty assistant response";
const TOOL_ERROR: &str = "WebSearch failed: unexpected tool call";
const STREAM_ERROR: &str = "WebSearch failed: response stream closed";
pub struct Request {
pub query: String,
pub model: String,
pub reasoning_effort: String,
pub deadline: Instant,
}
#[derive(Clone)]
pub struct Runner {
adapter: Adapter,
counter: Arc<AtomicU64>,
}
impl Runner {
pub fn new(adapter: Adapter) -> Self {
Self {
adapter,
counter: Arc::new(AtomicU64::new(0)),
}
}
pub async fn run(&self, request: Request) -> Result<String, String> {
let Request {
query,
model,
reasoning_effort,
deadline,
} = request;
if Instant::now() >= deadline {
return Err(DEADLINE_ERROR.into());
}
let key = next_key(&self.counter)?;
let start = self
.adapter
.start_web_search_turn(key, query, model, reasoning_effort);
let mut turn = match timeout_at(TokioInstant::from_std(deadline), start).await {
Ok(Ok(turn)) => turn,
Ok(Err(error)) => return Err(safe_runtime_error(&error)),
Err(_) => return Err(DEADLINE_ERROR.into()),
};
let mut text = String::new();
loop {
let event = match timeout_at(TokioInstant::from_std(deadline), turn.next_event()).await
{
Ok(event) => event,
Err(_) => return Err(DEADLINE_ERROR.into()),
};
match event {
Some(Event::TextDelta(delta)) => text.push_str(&delta),
Some(Event::Done) => {
if let Err(error) = require_message(&text) {
return fail_after_close(turn, deadline, error).await;
}
close_turn(turn, deadline).await?;
return Ok(text);
}
Some(Event::Error(error)) => {
let error = safe_runtime_error(&error);
return fail_after_close(turn, deadline, error).await;
}
Some(Event::ToolCall(_)) => {
return fail_after_close(turn, deadline, TOOL_ERROR.into()).await;
}
None => {
return fail_after_close(turn, deadline, STREAM_ERROR.into()).await;
}
}
}
}
}
fn next_key(counter: &AtomicU64) -> Result<String, String> {
let previous = counter
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |value| {
value.checked_add(1)
})
.map_err(|_| EXHAUSTED_ERROR.to_owned())?;
let number = previous
.checked_add(1)
.ok_or_else(|| EXHAUSTED_ERROR.to_owned())?;
Ok(format!("k1-websearch:{number:016}"))
}
fn require_message(text: &str) -> Result<(), String> {
(!text.is_empty())
.then_some(())
.ok_or_else(|| EMPTY_ERROR.to_owned())
}
async fn close_turn(turn: SearchTurn, deadline: Instant) -> Result<(), String> {
match timeout_at(TokioInstant::from_std(deadline), turn.close()).await {
Ok(Ok(())) => Ok(()),
Ok(Err(error)) => Err(safe_runtime_error(&error)),
Err(_) => Err(DEADLINE_ERROR.into()),
}
}
async fn fail_after_close(
turn: SearchTurn,
deadline: Instant,
error: String,
) -> Result<String, String> {
match close_turn(turn, deadline).await {
Ok(()) => Err(error),
Err(close_error) => Err(close_error),
}
}
fn safe_runtime_error(error: &Error) -> String {
match error.kind {
ErrorKind::Busy => "WebSearch failed: runtime busy",
ErrorKind::Interrupted => "WebSearch failed: interrupted",
ErrorKind::InvalidToolResult | ErrorKind::Protocol => {
"WebSearch failed: runtime protocol error"
}
ErrorKind::LaunchRejected => "WebSearch failed: launch rejected",
ErrorKind::Server => "WebSearch failed: provider error",
ErrorKind::Unavailable => "WebSearch failed: runtime unavailable",
}
.into()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn shared_counter_produces_exact_unique_keys_and_stops_at_exhaustion() {
let counter = Arc::new(AtomicU64::new(0));
let clone = Arc::clone(&counter);
assert_eq!(next_key(&counter).unwrap(), "k1-websearch:0000000000000001");
assert_eq!(next_key(&clone).unwrap(), "k1-websearch:0000000000000002");
counter.store(u64::MAX, Ordering::Relaxed);
assert_eq!(next_key(&counter).unwrap_err(), EXHAUSTED_ERROR);
assert_eq!(counter.load(Ordering::Relaxed), u64::MAX);
}
#[test]
fn runtime_errors_map_to_safe_stable_categories() {
for (kind, expected) in [
(ErrorKind::Busy, "WebSearch failed: runtime busy"),
(ErrorKind::Interrupted, "WebSearch failed: interrupted"),
(
ErrorKind::InvalidToolResult,
"WebSearch failed: runtime protocol error",
),
(
ErrorKind::LaunchRejected,
"WebSearch failed: launch rejected",
),
(
ErrorKind::Protocol,
"WebSearch failed: runtime protocol error",
),
(ErrorKind::Server, "WebSearch failed: provider error"),
(
ErrorKind::Unavailable,
"WebSearch failed: runtime unavailable",
),
] {
let error = Error::new(kind, "private detail").with_diagnostics(vec![1, 2, 3]);
assert_eq!(safe_runtime_error(&error), expected);
}
}
#[test]
fn completed_response_requires_only_nonempty_verbatim_text() {
assert_eq!(require_message("").unwrap_err(), EMPTY_ERROR);
assert!(require_message(" ").is_ok());
assert!(require_message("answer\nhttps://example.com").is_ok());
}
}