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}