1use std::collections::HashMap;
2
3use anyhow::{Context, Result};
4use async_openai::Client;
5use async_openai::config::OpenAIConfig;
6use async_openai::types::stream::StreamResponse;
7use async_trait::async_trait;
8use futures_util::StreamExt;
9use serde_json::Value;
10use tokio_util::sync::CancellationToken;
11
12use crate::config::ProviderConfig;
13use crate::model::{ContentBlock, ModelRequest, ModelTurn, Role, StreamEvent, ToolCall, Usage};
14
15use super::{EventSink, Provider, merge_request_fields, tool_definitions};
16
17pub struct ResponsesProvider {
18 client: Client<OpenAIConfig>,
19 config: ProviderConfig,
20}
21
22impl ResponsesProvider {
23 pub fn new(config: ProviderConfig, api_key: String) -> Result<Self> {
24 let base_url = config
25 .base_url
26 .clone()
27 .unwrap_or_else(|| "https://api.openai.com/v1".into());
28 let mut sdk_config = OpenAIConfig::new()
29 .with_api_key(api_key)
30 .with_api_base(base_url.trim_end_matches('/'));
31 for (key, value) in &config.headers {
32 let name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?;
33 sdk_config = sdk_config.with_header(name, value.as_str())?;
34 }
35 Ok(Self {
36 client: Client::with_config(sdk_config),
37 config,
38 })
39 }
40
41 fn request_body(&self, request: ModelRequest) -> Value {
42 let mut input = Vec::new();
43 for message in request.messages {
44 match message.role {
45 Role::User => {
46 let text = text_blocks(&message.blocks);
47 if !text.is_empty() {
48 input.push(serde_json::json!({"role":"user","content":text}));
49 }
50 }
51 Role::Assistant => {
52 let text = text_blocks(&message.blocks);
53 if !text.is_empty() {
54 input.push(serde_json::json!({"role":"assistant","content":text}));
55 }
56 for block in message.blocks {
57 if let ContentBlock::ToolCall(call) = block {
58 input.push(serde_json::json!({
59 "type":"function_call", "call_id":call.id,
60 "name":call.name, "arguments":call.arguments
61 }));
62 }
63 }
64 }
65 Role::Tool => {
66 for block in message.blocks {
67 if let ContentBlock::ToolResult(result) = block {
68 input.push(serde_json::json!({
69 "type":"function_call_output", "call_id":result.call_id,
70 "output":result.output
71 }));
72 }
73 }
74 }
75 Role::System => {}
76 }
77 }
78 let tools = if request.include_tools {
79 tool_definitions()
80 .into_iter()
81 .map(|mut tool| {
82 tool.as_object_mut()
83 .expect("tool definition is an object")
84 .insert("type".into(), Value::String("function".into()));
85 tool
86 })
87 .collect::<Vec<_>>()
88 } else {
89 Vec::new()
90 };
91 let mut body = serde_json::Map::new();
92 merge_request_fields(&mut body, &self.config);
93 body.insert("model".into(), Value::String(self.config.model.clone()));
94 body.insert(
95 "max_output_tokens".into(),
96 Value::from(self.config.max_tokens),
97 );
98 body.insert("instructions".into(), Value::String(request.system_prompt));
99 body.insert("input".into(), Value::Array(input));
100 body.insert("tools".into(), Value::Array(tools));
101 body.insert("stream".into(), Value::Bool(true));
102 Value::Object(body)
103 }
104}
105
106#[async_trait]
107impl Provider for ResponsesProvider {
108 async fn stream_turn(
109 &self,
110 request: ModelRequest,
111 events: EventSink,
112 cancel: CancellationToken,
113 ) -> Result<ModelTurn> {
114 let responses = self.client.responses();
115 let create = responses.create_stream_byot(self.request_body(request));
116 tokio::pin!(create);
117 let mut stream: StreamResponse<Value> = tokio::select! {
118 _ = cancel.cancelled() => anyhow::bail!("Responses API request cancelled"),
119 result = &mut create => result.context("start Responses API stream")?,
120 };
121 let mut values = Vec::new();
122 let mut live = ResponsesLive::default();
123 loop {
124 tokio::select! {
125 _ = cancel.cancelled() => anyhow::bail!("Responses API request cancelled"),
126 item = stream.next() => match item {
127 Some(Ok(value)) => {
128 live.emit(&value, &events);
129 values.push(value);
130 }
131 Some(Err(error)) => return Err(error).context("read Responses API stream"),
132 None => break,
133 }
134 }
135 }
136 normalize_events(values).map(|(turn, _)| turn)
137 }
138}
139
140#[derive(Default)]
141struct ResponsesLive {
142 calls: HashMap<String, String>,
143}
144
145impl ResponsesLive {
146 fn emit(&mut self, value: &Value, sink: &EventSink) {
147 match value
148 .get("type")
149 .and_then(Value::as_str)
150 .unwrap_or_default()
151 {
152 "response.output_text.delta" => sink.emit(StreamEvent::TextDelta {
153 delta: string(value, "delta"),
154 }),
155 "response.reasoning_summary_text.delta" | "response.reasoning_text.delta" => {
156 sink.emit(StreamEvent::ReasoningDelta {
157 delta: string(value, "delta"),
158 })
159 }
160 "response.output_item.added" if value["item"]["type"] == "function_call" => {
161 let item = &value["item"];
162 let item_id = string(item, "id");
163 let id = item
164 .get("call_id")
165 .and_then(Value::as_str)
166 .unwrap_or(&item_id)
167 .to_owned();
168 self.calls.insert(item_id, id.clone());
169 sink.emit(StreamEvent::ToolCallStart {
170 id: id.clone(),
171 name: string(item, "name"),
172 });
173 let arguments = string(item, "arguments");
174 if !arguments.is_empty() {
175 sink.emit(StreamEvent::ToolCallArgsDelta {
176 id,
177 delta: arguments,
178 });
179 }
180 }
181 "response.function_call_arguments.delta" => {
182 if let Some(id) = self.calls.get(&string(value, "item_id")) {
183 sink.emit(StreamEvent::ToolCallArgsDelta {
184 id: id.clone(),
185 delta: string(value, "delta"),
186 });
187 }
188 }
189 "response.output_item.done" if value["item"]["type"] == "function_call" => {
190 let item_id = string(&value["item"], "id");
191 if let Some(id) = self.calls.get(&item_id) {
192 sink.emit(StreamEvent::ToolCallEnd { id: id.clone() });
193 }
194 }
195 "response.completed" => {
196 sink.emit(StreamEvent::Usage(normalize_usage(
197 &value["response"]["usage"],
198 )));
199 sink.emit(StreamEvent::Done);
200 }
201 "error" => sink.emit(StreamEvent::Error {
202 message: value.to_string(),
203 }),
204 _ => {}
205 }
206 }
207}
208
209fn text_blocks(blocks: &[ContentBlock]) -> String {
210 blocks
211 .iter()
212 .filter_map(|block| match block {
213 ContentBlock::Text(text) => Some(text.as_str()),
214 _ => None,
215 })
216 .collect::<Vec<_>>()
217 .join("\n")
218}
219
220pub fn normalize_events(values: Vec<Value>) -> Result<(ModelTurn, Vec<StreamEvent>)> {
221 let mut text = String::new();
222 let mut reasoning = String::new();
223 let mut calls: HashMap<String, ToolCall> = HashMap::new();
224 let mut order = Vec::new();
225 let mut stream_events = Vec::new();
226 let mut usage = None;
227
228 for value in values {
229 match value
230 .get("type")
231 .and_then(Value::as_str)
232 .unwrap_or_default()
233 {
234 "response.output_text.delta" => {
235 let delta = string(&value, "delta");
236 text.push_str(&delta);
237 stream_events.push(StreamEvent::TextDelta { delta });
238 }
239 "response.reasoning_summary_text.delta" | "response.reasoning_text.delta" => {
240 let delta = string(&value, "delta");
241 reasoning.push_str(&delta);
242 stream_events.push(StreamEvent::ReasoningDelta { delta });
243 }
244 "response.output_item.added" => {
245 let item = &value["item"];
246 if item.get("type").and_then(Value::as_str) == Some("function_call") {
247 let item_id = string(item, "id");
248 let call = ToolCall::new(
249 item.get("call_id")
250 .and_then(Value::as_str)
251 .unwrap_or(&item_id),
252 string(item, "name"),
253 string(item, "arguments"),
254 );
255 stream_events.push(StreamEvent::ToolCallStart {
256 id: call.id.clone(),
257 name: call.name.clone(),
258 });
259 if !call.arguments.is_empty() {
260 stream_events.push(StreamEvent::ToolCallArgsDelta {
261 id: call.id.clone(),
262 delta: call.arguments.clone(),
263 });
264 }
265 order.push(item_id.clone());
266 calls.insert(item_id, call);
267 }
268 }
269 "response.function_call_arguments.delta" => {
270 let item_id = string(&value, "item_id");
271 let delta = string(&value, "delta");
272 let call = calls.get_mut(&item_id).with_context(|| {
273 format!("arguments for unknown function call item {item_id}")
274 })?;
275 call.arguments.push_str(&delta);
276 stream_events.push(StreamEvent::ToolCallArgsDelta {
277 id: call.id.clone(),
278 delta,
279 });
280 }
281 "response.output_item.done" => {
282 let item = &value["item"];
283 if item.get("type").and_then(Value::as_str) == Some("function_call") {
284 let item_id = string(item, "id");
285 let final_arguments = string(item, "arguments");
286 if let Some(call) = calls.get_mut(&item_id) {
287 if !final_arguments.is_empty() {
288 call.arguments = final_arguments;
289 }
290 stream_events.push(StreamEvent::ToolCallEnd {
291 id: call.id.clone(),
292 });
293 }
294 }
295 }
296 "response.completed" => {
297 usage = Some(normalize_usage(&value["response"]["usage"]));
298 }
299 "response.failed" | "response.incomplete" => {
300 anyhow::bail!("provider response did not complete: {}", value);
301 }
302 "error" => {
303 let message = string(&value, "message");
304 let code = string(&value, "code");
305 anyhow::bail!("provider error: {message} ({code})");
306 }
307 _ => {}
308 }
309 }
310 let tool_calls = order
311 .into_iter()
312 .filter_map(|id| calls.remove(&id))
313 .collect::<Vec<_>>();
314 let mut blocks = Vec::new();
315 if !reasoning.is_empty() {
316 blocks.push(ContentBlock::Reasoning(reasoning));
317 }
318 if !text.is_empty() {
319 blocks.push(ContentBlock::Text(text));
320 }
321 blocks.extend(tool_calls.iter().cloned().map(ContentBlock::ToolCall));
322 if let Some(usage) = usage {
323 stream_events.push(StreamEvent::Usage(usage));
324 }
325 stream_events.push(StreamEvent::Done);
326 Ok((
327 ModelTurn {
328 blocks,
329 tool_calls,
330 usage,
331 provider_state: None,
332 },
333 stream_events,
334 ))
335}
336
337fn normalize_usage(raw: &Value) -> Usage {
338 let cached_tokens = raw
339 .pointer("/input_tokens_details/cached_tokens")
340 .and_then(Value::as_u64);
341 let cache_write_tokens = raw
342 .pointer("/input_tokens_details/cache_write_tokens")
343 .and_then(Value::as_u64);
344 Usage {
345 input_tokens: raw
346 .get("input_tokens")
347 .and_then(Value::as_u64)
348 .map(|input| {
349 input.saturating_sub(cached_tokens.unwrap_or(0) + cache_write_tokens.unwrap_or(0))
350 }),
351 output_tokens: raw.get("output_tokens").and_then(Value::as_u64),
352 cached_tokens,
353 cache_write_tokens,
354 total_tokens: raw.get("total_tokens").and_then(Value::as_u64),
355 }
356}
357
358fn string(value: &Value, key: &str) -> String {
359 value
360 .get(key)
361 .and_then(Value::as_str)
362 .unwrap_or_default()
363 .to_owned()
364}