1use crate::error::AgentRuntimeError;
18use async_trait::async_trait;
19
20#[derive(Debug)]
28pub struct CompletionOptions<'a> {
29 pub model: &'a str,
31 pub max_tokens: Option<usize>,
33 pub temperature: Option<f32>,
35 pub timeout: Option<std::time::Duration>,
37 pub stop_sequences: Vec<String>,
40}
41
42impl<'a> CompletionOptions<'a> {
43 pub fn new(model: &'a str) -> Self {
45 Self {
46 model,
47 max_tokens: None,
48 temperature: None,
49 timeout: None,
50 stop_sequences: vec![],
51 }
52 }
53
54 pub fn with_max_tokens(mut self, n: usize) -> Self {
56 self.max_tokens = Some(n);
57 self
58 }
59
60 pub fn with_temperature(mut self, t: f32) -> Self {
62 self.temperature = Some(t);
63 self
64 }
65
66 pub fn with_timeout(mut self, d: std::time::Duration) -> Self {
68 self.timeout = Some(d);
69 self
70 }
71
72 pub fn with_stop_sequences(mut self, sequences: Vec<String>) -> Self {
74 self.stop_sequences = sequences;
75 self
76 }
77
78 pub fn with_timeout_secs(self, secs: u64) -> Self {
80 self.with_timeout(std::time::Duration::from_secs(secs))
81 }
82
83 pub fn with_timeout_ms(self, ms: u64) -> Self {
85 self.with_timeout(std::time::Duration::from_millis(ms))
86 }
87
88 pub fn has_stop_sequences(&self) -> bool {
90 !self.stop_sequences.is_empty()
91 }
92
93 pub fn stop_sequence_count(&self) -> usize {
95 self.stop_sequences.len()
96 }
97}
98
99#[async_trait]
107pub trait LlmProvider: Send + Sync {
108 async fn complete(&self, prompt: &str, model: &str) -> Result<String, AgentRuntimeError>;
114
115 async fn complete_with_options(
122 &self,
123 prompt: &str,
124 options: CompletionOptions<'_>,
125 ) -> Result<String, AgentRuntimeError> {
126 self.complete(prompt, options.model).await
127 }
128
129 async fn stream_complete(
149 &self,
150 prompt: &str,
151 model: &str,
152 ) -> Result<tokio::sync::mpsc::Receiver<Result<String, AgentRuntimeError>>, AgentRuntimeError>
153 {
154 let result = self.complete(prompt, model).await;
155 let (tx, rx) = tokio::sync::mpsc::channel(64);
156 let _ = tx.send(result).await;
158 Ok(rx)
159 }
160}
161
162#[cfg(feature = "anthropic")]
165pub struct AnthropicProvider {
175 api_key: String,
176 api_url: String,
179 client: reqwest::Client,
180 stream_semaphore: std::sync::Arc<tokio::sync::Semaphore>,
182 stream_max_tokens: Option<u32>,
185}
186
187#[cfg(feature = "anthropic")]
188impl AnthropicProvider {
189 const DEFAULT_API_URL: &'static str = "https://api.anthropic.com/v1/messages";
190 const API_VERSION: &'static str = "2023-06-01";
191 const MAX_TOKENS: u32 = 1024;
192 const DEFAULT_STREAM_CONCURRENCY: usize = 32;
194
195 pub fn new(api_key: impl Into<String>) -> Self {
197 Self {
198 api_key: api_key.into(),
199 api_url: Self::DEFAULT_API_URL.to_owned(),
200 client: reqwest::Client::new(),
201 stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
202 Self::DEFAULT_STREAM_CONCURRENCY,
203 )),
204 stream_max_tokens: None,
205 }
206 }
207
208 pub fn with_base_url(api_key: impl Into<String>, api_url: impl Into<String>) -> Self {
212 Self {
213 api_key: api_key.into(),
214 api_url: api_url.into(),
215 client: reqwest::Client::new(),
216 stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
217 Self::DEFAULT_STREAM_CONCURRENCY,
218 )),
219 stream_max_tokens: None,
220 }
221 }
222
223 pub fn with_max_concurrent_streams(api_key: impl Into<String>, max: usize) -> Self {
228 Self {
229 api_key: api_key.into(),
230 api_url: Self::DEFAULT_API_URL.to_owned(),
231 client: reqwest::Client::new(),
232 stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(max)),
233 stream_max_tokens: None,
234 }
235 }
236
237 pub fn with_stream_max_tokens(mut self, max_tokens: u32) -> Self {
245 self.stream_max_tokens = Some(max_tokens);
246 self
247 }
248}
249
250#[cfg(feature = "anthropic")]
251#[async_trait]
252impl LlmProvider for AnthropicProvider {
253 async fn complete(&self, prompt: &str, model: &str) -> Result<String, AgentRuntimeError> {
254 self.complete_with_options(prompt, CompletionOptions::new(model))
255 .await
256 }
257
258 #[tracing::instrument(skip(self, prompt, options), fields(model = options.model, provider = "anthropic"))]
262 async fn complete_with_options(
263 &self,
264 prompt: &str,
265 options: CompletionOptions<'_>,
266 ) -> Result<String, AgentRuntimeError> {
267 let max_tokens = options
268 .max_tokens
269 .unwrap_or(Self::MAX_TOKENS as usize) as u32;
270
271 let mut body = serde_json::json!({
272 "model": options.model,
273 "max_tokens": max_tokens,
274 "messages": [{ "role": "user", "content": prompt }]
275 });
276 if let Some(t) = options.temperature {
277 body["temperature"] = serde_json::json!(t);
278 }
279
280 let mut req = self
281 .client
282 .post(&self.api_url)
283 .header("x-api-key", &self.api_key)
284 .header("anthropic-version", Self::API_VERSION)
285 .header("content-type", "application/json")
286 .json(&body);
287 if let Some(timeout) = options.timeout {
288 req = req.timeout(timeout);
289 }
290 let response = req
291 .send()
292 .await
293 .map_err(|e| AgentRuntimeError::Provider(format!("Anthropic request failed: {e}")))?;
294
295 if !response.status().is_success() {
296 let status = response.status();
297 let text = response.text().await.unwrap_or_default();
298 return Err(AgentRuntimeError::Provider(format!(
299 "Anthropic API error {status}: {text}"
300 )));
301 }
302
303 let json: serde_json::Value = response
304 .json()
305 .await
306 .map_err(|e| AgentRuntimeError::Provider(format!("Anthropic parse failed: {e}")))?;
307
308 let text = json["content"]
309 .as_array()
310 .and_then(|arr| arr.first())
311 .and_then(|block| block["text"].as_str())
312 .ok_or_else(|| {
313 AgentRuntimeError::Provider("Anthropic response missing content[0].text".into())
314 })?;
315
316 Ok(text.to_owned())
317 }
318
319 async fn stream_complete(
325 &self,
326 prompt: &str,
327 model: &str,
328 ) -> Result<tokio::sync::mpsc::Receiver<Result<String, AgentRuntimeError>>, AgentRuntimeError>
329 {
330 let max_tokens = self.stream_max_tokens.unwrap_or(Self::MAX_TOKENS);
331 let body = serde_json::json!({
332 "model": model,
333 "max_tokens": max_tokens,
334 "stream": true,
335 "messages": [{ "role": "user", "content": prompt }]
336 });
337
338 let response = self
339 .client
340 .post(&self.api_url)
341 .header("x-api-key", &self.api_key)
342 .header("anthropic-version", Self::API_VERSION)
343 .header("content-type", "application/json")
344 .json(&body)
345 .send()
346 .await
347 .map_err(|e| {
348 AgentRuntimeError::Provider(format!("Anthropic stream request failed: {e}"))
349 })?;
350
351 if !response.status().is_success() {
352 let status = response.status();
353 let text = response.text().await.unwrap_or_default();
354 return Err(AgentRuntimeError::Provider(format!(
355 "Anthropic stream API error {status}: {text}"
356 )));
357 }
358
359 let (tx, rx) = tokio::sync::mpsc::channel::<Result<String, AgentRuntimeError>>(32);
360
361 let permit = std::sync::Arc::clone(&self.stream_semaphore)
362 .acquire_owned()
363 .await
364 .map_err(|_| {
365 AgentRuntimeError::Provider("Anthropic stream semaphore closed".into())
366 })?;
367
368 tokio::spawn(async move {
371 let _permit = permit;
372 let mut response = response;
373 let mut buffer = String::new();
374 loop {
375 match response.chunk().await {
376 Ok(Some(chunk)) => {
377 match String::from_utf8(chunk.to_vec()) {
378 Ok(s) => buffer.push_str(&s),
379 Err(e) => {
380 let _ = tx
381 .send(Err(AgentRuntimeError::Provider(format!(
382 "Anthropic stream: invalid UTF-8 in chunk: {e}"
383 ))))
384 .await;
385 return;
386 }
387 }
388 while let Some(newline) = buffer.find('\n') {
390 let line = buffer[..newline].trim().to_owned();
391 buffer = buffer[newline + 1..].to_owned();
392 if let Some(data) = line.strip_prefix("data: ") {
393 if data == "[DONE]" {
394 return;
395 }
396 if let Ok(json) =
397 serde_json::from_str::<serde_json::Value>(data)
398 {
399 if let Some(delta) = json["delta"]["text"].as_str() {
400 if tx.send(Ok(delta.to_owned())).await.is_err() {
401 return;
402 }
403 }
404 }
405 }
406 }
407 }
408 Ok(None) => break,
409 Err(e) => {
410 let _ = tx
411 .send(Err(AgentRuntimeError::Provider(format!(
412 "Anthropic stream chunk error: {e}"
413 ))))
414 .await;
415 return;
416 }
417 }
418 }
419 });
420
421 Ok(rx)
422 }
423}
424
425#[cfg(feature = "openai")]
428pub struct OpenAiProvider {
441 api_key: String,
442 base_url: String,
443 client: reqwest::Client,
444 stream_semaphore: std::sync::Arc<tokio::sync::Semaphore>,
446}
447
448#[cfg(feature = "openai")]
449impl OpenAiProvider {
450 const DEFAULT_BASE_URL: &'static str = "https://api.openai.com/v1";
451 const DEFAULT_STREAM_CONCURRENCY: usize = 32;
453
454 pub fn new(api_key: impl Into<String>) -> Self {
456 Self {
457 api_key: api_key.into(),
458 base_url: Self::DEFAULT_BASE_URL.to_owned(),
459 client: reqwest::Client::new(),
460 stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
461 Self::DEFAULT_STREAM_CONCURRENCY,
462 )),
463 }
464 }
465
466 pub fn with_base_url(api_key: impl Into<String>, base_url: impl Into<String>) -> Self {
468 Self {
469 api_key: api_key.into(),
470 base_url: base_url.into(),
471 client: reqwest::Client::new(),
472 stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
473 Self::DEFAULT_STREAM_CONCURRENCY,
474 )),
475 }
476 }
477
478 pub fn with_max_concurrent_streams(
480 api_key: impl Into<String>,
481 base_url: impl Into<String>,
482 max: usize,
483 ) -> Self {
484 Self {
485 api_key: api_key.into(),
486 base_url: base_url.into(),
487 client: reqwest::Client::new(),
488 stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(max)),
489 }
490 }
491}
492
493#[cfg(feature = "openai")]
494#[async_trait]
495impl LlmProvider for OpenAiProvider {
496 #[tracing::instrument(skip(self, prompt), fields(model, provider = "openai"))]
497 async fn complete(&self, prompt: &str, model: &str) -> Result<String, AgentRuntimeError> {
498 self.complete_with_options(prompt, CompletionOptions::new(model))
499 .await
500 }
501
502 #[tracing::instrument(skip(self, prompt, options), fields(model = options.model, provider = "openai"))]
507 async fn complete_with_options(
508 &self,
509 prompt: &str,
510 options: CompletionOptions<'_>,
511 ) -> Result<String, AgentRuntimeError> {
512 let url = format!("{}/chat/completions", self.base_url);
513 let mut body = serde_json::json!({
514 "model": options.model,
515 "messages": [{ "role": "user", "content": prompt }]
516 });
517 if let Some(max_tokens) = options.max_tokens {
518 body["max_tokens"] = serde_json::json!(max_tokens);
519 }
520 if let Some(temp) = options.temperature {
521 body["temperature"] = serde_json::json!(temp);
522 }
523
524 let mut req = self
525 .client
526 .post(&url)
527 .bearer_auth(&self.api_key)
528 .header("content-type", "application/json")
529 .json(&body);
530 if let Some(timeout) = options.timeout {
531 req = req.timeout(timeout);
532 }
533 let response = req
534 .send()
535 .await
536 .map_err(|e| AgentRuntimeError::Provider(format!("OpenAI request failed: {e}")))?;
537
538 if !response.status().is_success() {
539 let status = response.status();
540 let text = response.text().await.unwrap_or_default();
541 return Err(AgentRuntimeError::Provider(format!(
542 "OpenAI API error {status}: {text}"
543 )));
544 }
545
546 let json: serde_json::Value = response
547 .json()
548 .await
549 .map_err(|e| AgentRuntimeError::Provider(format!("OpenAI parse failed: {e}")))?;
550
551 let text = json["choices"]
552 .as_array()
553 .and_then(|arr| arr.first())
554 .and_then(|choice| choice["message"]["content"].as_str())
555 .ok_or_else(|| {
556 AgentRuntimeError::Provider(
557 "OpenAI response missing choices[0].message.content".into(),
558 )
559 })?;
560
561 Ok(text.to_owned())
562 }
563
564 async fn stream_complete(
565 &self,
566 prompt: &str,
567 model: &str,
568 ) -> Result<tokio::sync::mpsc::Receiver<Result<String, AgentRuntimeError>>, AgentRuntimeError>
569 {
570 let url = format!("{}/chat/completions", self.base_url);
571 let body = serde_json::json!({
572 "model": model,
573 "stream": true,
574 "messages": [{ "role": "user", "content": prompt }]
575 });
576
577 let mut response = self
578 .client
579 .post(&url)
580 .bearer_auth(&self.api_key)
581 .header("content-type", "application/json")
582 .json(&body)
583 .send()
584 .await
585 .map_err(|e| {
586 AgentRuntimeError::Provider(format!("OpenAI stream request failed: {e}"))
587 })?;
588
589 if !response.status().is_success() {
590 let status = response.status();
591 let text = response.text().await.unwrap_or_default();
592 return Err(AgentRuntimeError::Provider(format!(
593 "OpenAI stream API error {status}: {text}"
594 )));
595 }
596
597 let (tx, rx) = tokio::sync::mpsc::channel::<Result<String, AgentRuntimeError>>(32);
598
599 let permit = std::sync::Arc::clone(&self.stream_semaphore)
601 .acquire_owned()
602 .await
603 .map_err(|_| {
604 AgentRuntimeError::Provider("OpenAI stream semaphore closed".into())
605 })?;
606
607 tokio::spawn(async move {
608 let _permit = permit;
609 let mut buffer = String::new();
612 loop {
613 match response.chunk().await {
614 Ok(Some(chunk)) => {
615 let text = match std::str::from_utf8(&chunk) {
616 Ok(t) => t,
617 Err(e) => {
618 let _ = tx
619 .send(Err(AgentRuntimeError::Provider(format!(
620 "OpenAI stream chunk is not valid UTF-8: {e}"
621 ))))
622 .await;
623 return;
624 }
625 };
626 buffer.push_str(text);
627 while let Some(newline) = buffer.find('\n') {
629 let line: String = buffer.drain(..=newline).collect();
630 let line = line.trim_end_matches(['\r', '\n']);
631 if let Some(data) = line.strip_prefix("data: ") {
632 if data == "[DONE]" {
633 return;
634 }
635 if let Ok(json) =
636 serde_json::from_str::<serde_json::Value>(data)
637 {
638 if let Some(content) = json["choices"]
639 .as_array()
640 .and_then(|c| c.first())
641 .and_then(|c| c["delta"]["content"].as_str())
642 {
643 if tx.send(Ok(content.to_owned())).await.is_err() {
644 return;
645 }
646 }
647 }
648 }
649 }
650 }
651 Ok(None) => break,
652 Err(e) => {
653 let _ = tx
654 .send(Err(AgentRuntimeError::Provider(format!(
655 "OpenAI stream read failed: {e}"
656 ))))
657 .await;
658 return;
659 }
660 }
661 }
662 });
663
664 Ok(rx)
665 }
666}
667
668#[cfg(test)]
671mod tests {
672 use super::*;
673 use std::sync::Arc;
674
675 struct StubProvider {
677 response: String,
678 }
679
680 #[async_trait]
681 impl LlmProvider for StubProvider {
682 async fn complete(&self, _prompt: &str, _model: &str) -> Result<String, AgentRuntimeError> {
683 Ok(self.response.clone())
684 }
685 }
687
688 #[tokio::test]
689 async fn test_stub_provider_returns_configured_response() {
690 let p = StubProvider {
691 response: "hello".into(),
692 };
693 let result = p.complete("prompt", "stub-model").await.unwrap();
694 assert_eq!(result, "hello");
695 }
696
697 #[tokio::test]
698 async fn test_llm_provider_is_object_safe() {
699 let p: Arc<dyn LlmProvider> = Arc::new(StubProvider {
700 response: "ok".into(),
701 });
702 let result = p.complete("test", "model").await.unwrap();
703 assert_eq!(result, "ok");
704 }
705
706 #[tokio::test]
707 async fn test_stub_provider_ignores_model_parameter() {
708 let p = StubProvider {
709 response: "42".into(),
710 };
711 let r1 = p.complete("q", "model-a").await.unwrap();
712 let r2 = p.complete("q", "model-b").await.unwrap();
713 assert_eq!(r1, r2);
714 }
715
716 #[tokio::test]
717 async fn test_stub_provider_stream_returns_single_chunk() {
718 let p = StubProvider {
719 response: "hello world".into(),
720 };
721 let mut rx = p.stream_complete("prompt", "model").await.unwrap();
722 let mut collected = String::new();
723 while let Some(chunk) = rx.recv().await {
724 collected.push_str(&chunk.unwrap());
725 }
726 assert_eq!(collected, "hello world");
727 }
728
729 #[tokio::test]
730 async fn test_stream_receiver_closes_after_completion() {
731 let p = StubProvider {
732 response: "done".into(),
733 };
734 let mut rx = p.stream_complete("prompt", "model").await.unwrap();
735 while let Some(_chunk) = rx.recv().await {}
737 assert!(rx.recv().await.is_none());
739 }
740}