Skip to main content

a_agent/provider/
anthropic.rs

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