Skip to main content

a_agent/provider/
responses.rs

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}