1use crate::{metrics, OrchestratorError};
34use async_trait::async_trait;
35use futures::Stream;
36use serde::{Deserialize, Serialize};
37use std::pin::Pin;
38use std::sync::Arc;
39use std::time::{Duration, Instant};
40
41pub type TokenStream =
43 Pin<Box<dyn Stream<Item = Result<String, OrchestratorError>> + Send + 'static>>;
44
45fn parse_retry_after(headers: &reqwest::header::HeaderMap) -> Option<Duration> {
48 let value = headers.get("retry-after")?.to_str().ok()?;
49 let secs: u64 = value.trim().parse().unwrap_or(60);
51 Some(Duration::from_secs(secs))
52}
53
54fn warn_if_low_remaining(headers: &reqwest::header::HeaderMap, provider: &str) {
56 if let Some(val) = headers
57 .get("x-ratelimit-remaining-requests")
58 .and_then(|v| v.to_str().ok())
59 .and_then(|s| s.parse::<u64>().ok())
60 {
61 if val < 10 {
62 tracing::warn!(
63 provider = provider,
64 remaining_requests = val,
65 "approaching provider rate limit"
66 );
67 }
68 }
69}
70
71#[async_trait]
83pub trait ModelWorker: Send + Sync {
84 async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError>;
91
92 async fn infer_stream(&self, prompt: &str) -> Result<TokenStream, OrchestratorError> {
98 let tokens = self.infer(prompt).await?;
99 let stream = futures::stream::iter(tokens.into_iter().map(Ok));
100 Ok(Box::pin(stream))
101 }
102}
103
104pub fn stream_worker(
128 worker: Arc<dyn ModelWorker>,
129 prompt: String,
130) -> tokio::sync::mpsc::Receiver<Result<String, OrchestratorError>> {
131 let (tx, rx) = tokio::sync::mpsc::channel(64);
132 let _stream_task = tokio::spawn(async move {
134 match worker.infer(&prompt).await {
135 Ok(tokens) => {
136 for token in tokens {
137 if tx.send(Ok(token)).await.is_err() {
138 break;
140 }
141 }
142 }
143 Err(e) => {
144 tracing::error!(error = %e, "stream_worker: inference failed");
145 let _ = tx.send(Err(e)).await;
146 }
147 }
148 });
149 rx
150}
151
152pub struct EchoWorker {
161 pub delay_ms: u64,
163}
164
165impl EchoWorker {
166 pub fn new() -> Self {
188 Self { delay_ms: 10 }
189 }
190
191 pub fn with_delay(delay_ms: u64) -> Self {
193 Self { delay_ms }
194 }
195}
196
197impl Default for EchoWorker {
198 fn default() -> Self {
199 Self::new()
200 }
201}
202
203#[async_trait]
204impl ModelWorker for EchoWorker {
205 async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError> {
206 tokio::time::sleep(tokio::time::Duration::from_millis(self.delay_ms)).await;
208
209 let tokens: Vec<String> = prompt.split_whitespace().map(str::to_string).collect();
211
212 Ok(tokens)
213 }
214}
215
216#[derive(Debug, Serialize)]
222struct OpenAiRequest {
223 model: String,
224 messages: Vec<OpenAiMessage>,
225 max_tokens: u32,
226 temperature: f32,
227}
228
229#[derive(Debug, Serialize)]
230struct OpenAiMessage {
231 role: String,
232 content: String,
233}
234
235#[derive(Debug, Deserialize)]
237struct OpenAiResponse {
238 choices: Vec<OpenAiChoice>,
239}
240
241#[derive(Debug, Deserialize)]
242struct OpenAiChoice {
243 message: OpenAiResponseMessage,
244}
245
246#[derive(Debug, Deserialize)]
247struct OpenAiResponseMessage {
248 content: String,
249}
250
251#[derive(Debug)]
274pub struct OpenAiWorker {
275 client: reqwest::Client,
276 api_key: String,
277 model: String,
278 max_tokens: u32,
279 temperature: f32,
280 timeout: Duration,
281 base_url: String,
283}
284
285impl OpenAiWorker {
286 pub fn new(model: impl Into<String>) -> Result<Self, OrchestratorError> {
319 let api_key = std::env::var("OPENAI_API_KEY").map_err(|_| {
320 OrchestratorError::ConfigError("OPENAI_API_KEY environment variable not set".into())
321 })?;
322
323 Ok(Self {
324 client: reqwest::Client::new(),
325 api_key,
326 model: model.into(),
327 max_tokens: 256,
328 temperature: 0.7,
329 timeout: Duration::from_secs(30),
330 base_url: "https://api.openai.com/v1".to_string(),
331 })
332 }
333
334 pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
336 self.max_tokens = max_tokens;
337 self
338 }
339
340 pub fn with_temperature(mut self, temperature: f32) -> Self {
345 if !(0.0..=2.0).contains(&temperature) {
346 tracing::warn!(
347 temperature = temperature,
348 "OpenAI temperature out of range [0.0, 2.0] — clamping"
349 );
350 self.temperature = temperature.clamp(0.0, 2.0);
351 } else {
352 self.temperature = temperature;
353 }
354 self
355 }
356
357 pub fn with_timeout(mut self, timeout: Duration) -> Self {
359 self.timeout = timeout;
360 self
361 }
362
363 pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
369 self.base_url = url.into();
370 self
371 }
372}
373
374#[async_trait]
375impl ModelWorker for OpenAiWorker {
376 async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError> {
377 let _infer_start = Instant::now();
378 let request = OpenAiRequest {
379 model: self.model.clone(),
380 messages: vec![OpenAiMessage {
381 role: "user".to_string(),
382 content: prompt.to_string(),
383 }],
384 max_tokens: self.max_tokens,
385 temperature: self.temperature,
386 };
387
388 let response = self
389 .client
390 .post(format!("{}/chat/completions", self.base_url))
391 .header("Authorization", format!("Bearer {}", self.api_key))
392 .header("Content-Type", "application/json")
393 .timeout(self.timeout)
394 .json(&request)
395 .send()
396 .await
397 .map_err(|e| OrchestratorError::Inference(format!("OpenAI request failed: {}", e)))?;
398
399 warn_if_low_remaining(response.headers(), "openai");
400
401 if response.status() == reqwest::StatusCode::TOO_MANY_REQUESTS {
402 let retry_after_secs =
403 parse_retry_after(response.headers()).unwrap_or(Duration::from_secs(60));
404 return Err(OrchestratorError::RateLimited {
405 retry_after_secs: retry_after_secs.as_secs(),
406 });
407 }
408
409 if !response.status().is_success() {
410 let status = response.status();
411 let error_text = response.text().await.unwrap_or_else(|_| String::new());
412 if status == reqwest::StatusCode::UNAUTHORIZED
413 || status == reqwest::StatusCode::FORBIDDEN
414 {
415 return Err(OrchestratorError::AuthFailed(format!("HTTP {status}")));
416 }
417 return Err(OrchestratorError::Inference(format!(
418 "OpenAI API error {}: {}",
419 status, error_text
420 )));
421 }
422
423 let api_response: OpenAiResponse = response.json().await.map_err(|e| {
424 OrchestratorError::Inference(format!("Failed to parse response: {}", e))
425 })?;
426
427 if api_response.choices.is_empty() {
428 return Err(OrchestratorError::Inference(
429 "No choices in OpenAI response".to_string(),
430 ));
431 }
432
433 let content = api_response
438 .choices
439 .first()
440 .ok_or_else(|| {
441 OrchestratorError::Inference("No choices in OpenAI response".to_string())
442 })?
443 .message
444 .content
445 .clone();
446
447 let result = if content.trim().is_empty() {
448 Ok(vec![])
449 } else {
450 Ok(vec![content])
451 };
452 tracing::debug!(
453 worker = "openai",
454 model = %self.model,
455 latency_ms = %_infer_start.elapsed().as_millis(),
456 "inference completed"
457 );
458 result
459 }
460
461 async fn infer_stream(&self, prompt: &str) -> Result<TokenStream, OrchestratorError> {
462 use futures::StreamExt;
463
464 #[derive(Deserialize)]
466 struct StreamChunk {
467 choices: Vec<StreamChoice>,
468 }
469 #[derive(Deserialize)]
470 struct StreamChoice {
471 delta: StreamDelta,
472 }
473 #[derive(Deserialize)]
474 struct StreamDelta {
475 #[serde(default)]
476 content: Option<String>,
477 }
478
479 let request = serde_json::json!({
480 "model": self.model,
481 "messages": [{"role": "user", "content": prompt}],
482 "max_tokens": self.max_tokens,
483 "temperature": self.temperature,
484 "stream": true
485 });
486
487 let response = self
488 .client
489 .post(format!("{}/chat/completions", self.base_url))
490 .header("Authorization", format!("Bearer {}", self.api_key))
491 .header("Content-Type", "application/json")
492 .timeout(self.timeout)
493 .json(&request)
494 .send()
495 .await
496 .map_err(|e| {
497 OrchestratorError::Inference(format!("OpenAI stream request failed: {e}"))
498 })?;
499
500 warn_if_low_remaining(response.headers(), "openai");
501
502 if response.status() == reqwest::StatusCode::TOO_MANY_REQUESTS {
503 let retry_after_secs =
504 parse_retry_after(response.headers()).unwrap_or(Duration::from_secs(60));
505 return Err(OrchestratorError::RateLimited {
506 retry_after_secs: retry_after_secs.as_secs(),
507 });
508 }
509
510 if !response.status().is_success() {
511 let status = response.status();
512 let body = response.text().await.unwrap_or_default();
513 if status == reqwest::StatusCode::UNAUTHORIZED
514 || status == reqwest::StatusCode::FORBIDDEN
515 {
516 return Err(OrchestratorError::AuthFailed(format!("HTTP {status}")));
517 }
518 return Err(OrchestratorError::Inference(format!(
519 "OpenAI stream error {status}: {body}"
520 )));
521 }
522
523 let request_start = Instant::now();
525 let model_name = self.model.clone();
526 let byte_stream = response.bytes_stream();
527 let token_stream = byte_stream.filter_map(|chunk| async move {
528 let bytes = chunk.ok()?;
529 let text = std::str::from_utf8(&bytes).ok()?;
530 let mut tokens = Vec::new();
532 for line in text.lines() {
533 let Some(json_str) = line.strip_prefix("data: ") else {
534 continue;
535 };
536 if json_str.trim() == "[DONE]" {
537 break;
538 }
539 if let Ok(chunk) = serde_json::from_str::<StreamChunk>(json_str) {
540 for choice in chunk.choices {
541 if let Some(content) = choice.delta.content {
542 if !content.is_empty() {
543 tokens.push(content);
544 }
545 }
546 }
547 }
548 }
549 if tokens.is_empty() {
550 None
551 } else {
552 Some(Ok(tokens.join("")))
553 }
554 });
555
556 let mut first_token_seen = false;
558 let ttft_stream = token_stream.map(move |item| {
559 if !first_token_seen {
560 first_token_seen = true;
561 metrics::record_ttft("openai", &model_name, request_start.elapsed());
562 }
563 item
564 });
565
566 Ok(Box::pin(ttft_stream))
567 }
568}
569
570#[derive(Debug)]
597pub struct AnthropicWorker {
598 client: reqwest::Client,
599 api_key: String,
600 model: String,
601 max_tokens: u32,
602 temperature: f32,
603 timeout: Duration,
604 base_url: String,
606}
607
608impl AnthropicWorker {
609 pub fn new(model: impl Into<String>) -> Result<Self, OrchestratorError> {
642 let api_key = std::env::var("ANTHROPIC_API_KEY").map_err(|_| {
643 OrchestratorError::ConfigError("ANTHROPIC_API_KEY environment variable not set".into())
644 })?;
645
646 Ok(Self {
647 client: reqwest::Client::new(),
648 api_key,
649 model: model.into(),
650 max_tokens: 1024,
651 temperature: 1.0,
652 timeout: Duration::from_secs(60),
653 base_url: "https://api.anthropic.com/v1".to_string(),
654 })
655 }
656
657 pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
659 self.max_tokens = max_tokens;
660 self
661 }
662
663 pub fn with_temperature(mut self, temperature: f32) -> Self {
668 if !(0.0..=1.0).contains(&temperature) {
669 tracing::warn!(
670 temperature = temperature,
671 "Anthropic temperature out of range [0.0, 1.0] — clamping"
672 );
673 self.temperature = temperature.clamp(0.0, 1.0);
674 } else {
675 self.temperature = temperature;
676 }
677 self
678 }
679
680 pub fn with_timeout(mut self, timeout: Duration) -> Self {
682 self.timeout = timeout;
683 self
684 }
685
686 pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
691 self.base_url = url.into();
692 self
693 }
694}
695
696#[async_trait]
697impl ModelWorker for AnthropicWorker {
698 async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError> {
699 let _infer_start = Instant::now();
700 let request = serde_json::json!({
702 "model": self.model,
703 "max_tokens": self.max_tokens,
704 "temperature": self.temperature,
705 "messages": [{"role": "user", "content": prompt}]
706 });
707
708 let response = self
709 .client
710 .post(format!("{}/messages", self.base_url))
711 .header("x-api-key", &self.api_key)
712 .header("anthropic-version", "2023-06-01")
713 .header("Content-Type", "application/json")
714 .timeout(self.timeout)
715 .json(&request)
716 .send()
717 .await
718 .map_err(|e| {
719 OrchestratorError::Inference(format!("Anthropic request failed: {}", e))
720 })?;
721
722 warn_if_low_remaining(response.headers(), "anthropic");
723
724 if response.status() == reqwest::StatusCode::TOO_MANY_REQUESTS {
725 let retry_after_secs =
726 parse_retry_after(response.headers()).unwrap_or(Duration::from_secs(60));
727 return Err(OrchestratorError::RateLimited {
728 retry_after_secs: retry_after_secs.as_secs(),
729 });
730 }
731
732 if !response.status().is_success() {
733 let status = response.status();
734 let error_text = response.text().await.unwrap_or_else(|_| String::new());
735 if status == reqwest::StatusCode::UNAUTHORIZED
736 || status == reqwest::StatusCode::FORBIDDEN
737 {
738 return Err(OrchestratorError::AuthFailed(format!("HTTP {status}")));
739 }
740 return Err(OrchestratorError::Inference(format!(
741 "Anthropic API error {}: {}",
742 status, error_text
743 )));
744 }
745
746 #[derive(Deserialize)]
748 struct MessagesResponse {
749 content: Vec<ContentBlock>,
750 }
751 #[derive(Deserialize)]
752 struct ContentBlock {
753 #[serde(rename = "type")]
754 block_type: String,
755 #[serde(default)]
756 text: String,
757 }
758
759 let api_response: MessagesResponse = response.json().await.map_err(|e| {
760 OrchestratorError::Inference(format!("Failed to parse response: {}", e))
761 })?;
762
763 let full_text: String = api_response
767 .content
768 .into_iter()
769 .filter(|b| b.block_type == "text")
770 .map(|b| b.text)
771 .collect::<Vec<_>>()
772 .join("");
773
774 let result = if full_text.trim().is_empty() {
775 Ok(vec![])
776 } else {
777 Ok(vec![full_text])
778 };
779 tracing::debug!(
780 worker = "anthropic",
781 model = %self.model,
782 latency_ms = %_infer_start.elapsed().as_millis(),
783 "inference completed"
784 );
785 result
786 }
787
788 async fn infer_stream(&self, prompt: &str) -> Result<TokenStream, OrchestratorError> {
789 use futures::StreamExt;
790
791 #[derive(Deserialize)]
794 struct StreamEvent {
795 #[serde(rename = "type")]
796 event_type: String,
797 delta: Option<StreamDelta>,
798 }
799 #[derive(Deserialize)]
800 struct StreamDelta {
801 #[serde(rename = "type")]
802 delta_type: String,
803 #[serde(default)]
804 text: String,
805 }
806
807 let request = serde_json::json!({
808 "model": self.model,
809 "max_tokens": self.max_tokens,
810 "temperature": self.temperature,
811 "stream": true,
812 "messages": [{"role": "user", "content": prompt}]
813 });
814
815 let response = self
816 .client
817 .post(format!("{}/messages", self.base_url))
818 .header("x-api-key", &self.api_key)
819 .header("anthropic-version", "2023-06-01")
820 .header("Content-Type", "application/json")
821 .timeout(self.timeout)
822 .json(&request)
823 .send()
824 .await
825 .map_err(|e| {
826 OrchestratorError::Inference(format!("Anthropic stream request failed: {e}"))
827 })?;
828
829 warn_if_low_remaining(response.headers(), "anthropic");
830
831 if response.status() == reqwest::StatusCode::TOO_MANY_REQUESTS {
832 let retry_after_secs =
833 parse_retry_after(response.headers()).unwrap_or(Duration::from_secs(60));
834 return Err(OrchestratorError::RateLimited {
835 retry_after_secs: retry_after_secs.as_secs(),
836 });
837 }
838
839 if !response.status().is_success() {
840 let status = response.status();
841 let body = response.text().await.unwrap_or_default();
842 if status == reqwest::StatusCode::UNAUTHORIZED
843 || status == reqwest::StatusCode::FORBIDDEN
844 {
845 return Err(OrchestratorError::AuthFailed(format!("HTTP {status}")));
846 }
847 return Err(OrchestratorError::Inference(format!(
848 "Anthropic stream error {status}: {body}"
849 )));
850 }
851
852 let request_start = Instant::now();
853 let model_name = self.model.clone();
854 let byte_stream = response.bytes_stream();
855 let token_stream = byte_stream.filter_map(|chunk| async move {
856 let bytes = chunk.ok()?;
857 let text = std::str::from_utf8(&bytes).ok()?;
858 let mut out = String::new();
859 for line in text.lines() {
860 let Some(json_str) = line.strip_prefix("data: ") else {
862 continue;
863 };
864 if json_str.trim() == "[DONE]" {
865 break;
866 }
867 if let Ok(event) = serde_json::from_str::<StreamEvent>(json_str) {
868 if event.event_type == "content_block_delta" {
869 if let Some(delta) = event.delta {
870 if delta.delta_type == "text_delta" && !delta.text.is_empty() {
871 out.push_str(&delta.text);
872 }
873 }
874 }
875 }
876 }
877 if out.is_empty() {
878 None
879 } else {
880 Some(Ok(out))
881 }
882 });
883
884 let mut first_token_seen = false;
886 let ttft_stream = token_stream.map(move |item| {
887 if !first_token_seen {
888 first_token_seen = true;
889 metrics::record_ttft("anthropic", &model_name, request_start.elapsed());
890 }
891 item
892 });
893
894 Ok(Box::pin(ttft_stream))
895 }
896}
897
898#[derive(Debug, Serialize)]
904struct LlamaCppRequest {
905 prompt: String,
906 n_predict: i32,
907 temperature: f32,
908 stop: Vec<String>,
909}
910
911#[derive(Debug, Deserialize)]
913struct LlamaCppResponse {
914 content: String,
915}
916
917pub struct LlamaCppWorker {
936 client: reqwest::Client,
937 url: String,
938 max_tokens: i32,
939 temperature: f32,
940 timeout: Duration,
941}
942
943impl LlamaCppWorker {
944 pub fn new() -> Self {
976 let url =
977 std::env::var("LLAMA_CPP_URL").unwrap_or_else(|_| "http://localhost:8080".to_string());
978
979 Self {
980 client: reqwest::Client::new(),
981 url,
982 max_tokens: 256,
983 temperature: 0.8,
984 timeout: Duration::from_secs(30),
985 }
986 }
987
988 pub fn with_url(mut self, url: impl Into<String>) -> Self {
990 self.url = url.into();
991 self
992 }
993
994 pub fn with_max_tokens(mut self, max_tokens: i32) -> Self {
996 self.max_tokens = max_tokens;
997 self
998 }
999
1000 pub fn with_temperature(mut self, temperature: f32) -> Self {
1005 if !(0.0..=2.0).contains(&temperature) {
1006 tracing::warn!(
1007 temperature = temperature,
1008 "LlamaCpp temperature out of range [0.0, 2.0] — clamping"
1009 );
1010 self.temperature = temperature.clamp(0.0, 2.0);
1011 } else {
1012 self.temperature = temperature;
1013 }
1014 self
1015 }
1016
1017 pub fn with_timeout(mut self, timeout: Duration) -> Self {
1019 self.timeout = timeout;
1020 self
1021 }
1022}
1023
1024impl Default for LlamaCppWorker {
1025 fn default() -> Self {
1026 Self::new()
1027 }
1028}
1029
1030#[async_trait]
1031impl ModelWorker for LlamaCppWorker {
1032 async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError> {
1033 let _infer_start = Instant::now();
1034 let request = LlamaCppRequest {
1035 prompt: prompt.to_string(),
1036 n_predict: self.max_tokens,
1037 temperature: self.temperature,
1038 stop: vec!["</s>".to_string(), "Human:".to_string()],
1039 };
1040
1041 let response = self
1042 .client
1043 .post(format!("{}/completion", self.url))
1044 .timeout(self.timeout)
1045 .json(&request)
1046 .send()
1047 .await
1048 .map_err(|e| {
1049 OrchestratorError::Inference(format!("llama.cpp request failed: {}", e))
1050 })?;
1051
1052 if !response.status().is_success() {
1053 let status = response.status();
1054 let error_text = response.text().await.unwrap_or_else(|_| String::new());
1055 if status == reqwest::StatusCode::UNAUTHORIZED
1056 || status == reqwest::StatusCode::FORBIDDEN
1057 {
1058 return Err(OrchestratorError::AuthFailed(format!("HTTP {status}")));
1059 }
1060 return Err(OrchestratorError::Inference(format!(
1061 "llama.cpp error {}: {}",
1062 status, error_text
1063 )));
1064 }
1065
1066 let api_response: LlamaCppResponse = response.json().await.map_err(|e| {
1067 OrchestratorError::Inference(format!("Failed to parse response: {}", e))
1068 })?;
1069
1070 let content = api_response.content;
1073 let result = if content.is_empty() {
1074 Ok(vec![])
1075 } else {
1076 Ok(vec![content])
1077 };
1078 tracing::debug!(
1079 worker = "llama_cpp",
1080 model = "llama.cpp",
1081 latency_ms = %_infer_start.elapsed().as_millis(),
1082 "inference completed"
1083 );
1084 result
1085 }
1086}
1087
1088#[derive(Debug, Serialize)]
1094struct VllmRequest {
1095 prompt: String,
1096 max_tokens: u32,
1097 temperature: f32,
1098 top_p: f32,
1099}
1100
1101#[derive(Debug, Deserialize)]
1103struct VllmResponse {
1104 text: Vec<String>,
1105}
1106
1107pub struct VllmWorker {
1126 client: reqwest::Client,
1127 url: String,
1128 max_tokens: u32,
1129 temperature: f32,
1130 top_p: f32,
1131 timeout: Duration,
1132}
1133
1134impl VllmWorker {
1135 pub fn new() -> Self {
1167 let url = std::env::var("VLLM_URL").unwrap_or_else(|_| "http://localhost:8000".to_string());
1168
1169 Self {
1170 client: reqwest::Client::new(),
1171 url,
1172 max_tokens: 512,
1173 temperature: 0.7,
1174 top_p: 0.95,
1175 timeout: Duration::from_secs(60),
1176 }
1177 }
1178
1179 pub fn with_url(mut self, url: impl Into<String>) -> Self {
1181 self.url = url.into();
1182 self
1183 }
1184
1185 pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
1187 self.max_tokens = max_tokens;
1188 self
1189 }
1190
1191 pub fn with_temperature(mut self, temperature: f32) -> Self {
1196 if !(0.0..=2.0).contains(&temperature) {
1197 tracing::warn!(
1198 temperature = temperature,
1199 "vLLM temperature out of range [0.0, 2.0] — clamping"
1200 );
1201 self.temperature = temperature.clamp(0.0, 2.0);
1202 } else {
1203 self.temperature = temperature;
1204 }
1205 self
1206 }
1207
1208 pub fn with_top_p(mut self, top_p: f32) -> Self {
1210 self.top_p = top_p;
1211 self
1212 }
1213
1214 pub fn with_timeout(mut self, timeout: Duration) -> Self {
1216 self.timeout = timeout;
1217 self
1218 }
1219}
1220
1221impl Default for VllmWorker {
1222 fn default() -> Self {
1223 Self::new()
1224 }
1225}
1226
1227#[async_trait]
1228impl ModelWorker for VllmWorker {
1229 async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError> {
1230 let _infer_start = Instant::now();
1231 let request = VllmRequest {
1232 prompt: prompt.to_string(),
1233 max_tokens: self.max_tokens,
1234 temperature: self.temperature,
1235 top_p: self.top_p,
1236 };
1237
1238 let response = self
1239 .client
1240 .post(format!("{}/generate", self.url))
1241 .timeout(self.timeout)
1242 .json(&request)
1243 .send()
1244 .await
1245 .map_err(|e| OrchestratorError::Inference(format!("vLLM request failed: {}", e)))?;
1246
1247 if !response.status().is_success() {
1248 let status = response.status();
1249 let error_text = response.text().await.unwrap_or_else(|_| String::new());
1250 if status == reqwest::StatusCode::UNAUTHORIZED
1251 || status == reqwest::StatusCode::FORBIDDEN
1252 {
1253 return Err(OrchestratorError::AuthFailed(format!("HTTP {status}")));
1254 }
1255 return Err(OrchestratorError::Inference(format!(
1256 "vLLM error {}: {}",
1257 status, error_text
1258 )));
1259 }
1260
1261 let api_response: VllmResponse = response.json().await.map_err(|e| {
1262 OrchestratorError::Inference(format!("Failed to parse response: {}", e))
1263 })?;
1264
1265 let content =
1266 api_response.text.into_iter().next().ok_or_else(|| {
1267 OrchestratorError::Inference("Empty response from vLLM".to_string())
1268 })?;
1269
1270 tracing::debug!(
1271 worker = "vllm",
1272 model = "vllm",
1273 latency_ms = %_infer_start.elapsed().as_millis(),
1274 "inference completed"
1275 );
1276 Ok(vec![content])
1277 }
1278}
1279
1280#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1286pub enum LoadBalanceStrategy {
1287 RoundRobin,
1289 LeastLoaded,
1291}
1292
1293pub struct LoadBalancedWorker {
1314 workers: Vec<Arc<dyn ModelWorker>>,
1315 strategy: LoadBalanceStrategy,
1316 rr_counter: std::sync::atomic::AtomicUsize,
1318 in_flight: Vec<std::sync::atomic::AtomicUsize>,
1320 names: Vec<String>,
1322}
1323
1324struct InFlightGuard<'a>(&'a std::sync::atomic::AtomicUsize);
1326impl Drop for InFlightGuard<'_> {
1327 fn drop(&mut self) {
1328 self.0.fetch_sub(1, std::sync::atomic::Ordering::Relaxed);
1329 }
1330}
1331
1332impl LoadBalancedWorker {
1333 pub fn round_robin(workers: Vec<Arc<dyn ModelWorker>>) -> Self {
1339 assert!(!workers.is_empty(), "worker pool must not be empty");
1340 let n = workers.len();
1341 Self {
1342 workers,
1343 strategy: LoadBalanceStrategy::RoundRobin,
1344 rr_counter: std::sync::atomic::AtomicUsize::new(0),
1345 in_flight: (0..n)
1346 .map(|_| std::sync::atomic::AtomicUsize::new(0))
1347 .collect(),
1348 names: Vec::new(),
1349 }
1350 }
1351
1352 pub fn replicate(worker: Arc<dyn ModelWorker>, n: usize) -> Self {
1361 assert!(n > 0, "replicate count must be > 0");
1362 let workers = std::iter::repeat_n(Arc::clone(&worker), n).collect();
1363 Self::round_robin(workers)
1364 }
1365
1366 pub fn least_loaded(workers: Vec<Arc<dyn ModelWorker>>) -> Self {
1372 assert!(!workers.is_empty(), "worker pool must not be empty");
1373 let n = workers.len();
1374 Self {
1375 workers,
1376 strategy: LoadBalanceStrategy::LeastLoaded,
1377 rr_counter: std::sync::atomic::AtomicUsize::new(0),
1378 in_flight: (0..n)
1379 .map(|_| std::sync::atomic::AtomicUsize::new(0))
1380 .collect(),
1381 names: Vec::new(),
1382 }
1383 }
1384
1385 pub fn with_names(mut self, names: Vec<String>) -> Self {
1389 self.names = names;
1390 self
1391 }
1392
1393 pub fn len(&self) -> usize {
1395 self.workers.len()
1396 }
1397
1398 pub fn is_empty(&self) -> bool {
1400 self.workers.is_empty()
1401 }
1402
1403 fn pick(&self) -> usize {
1405 match self.strategy {
1406 LoadBalanceStrategy::RoundRobin => {
1407 self.rr_counter
1408 .fetch_add(1, std::sync::atomic::Ordering::Relaxed)
1409 % self.workers.len()
1410 }
1411 LoadBalanceStrategy::LeastLoaded => self
1412 .in_flight
1413 .iter()
1414 .enumerate()
1415 .min_by_key(|(_, c)| c.load(std::sync::atomic::Ordering::Relaxed))
1416 .map(|(i, _)| i)
1417 .unwrap_or(0),
1418 }
1419 }
1420}
1421
1422#[async_trait]
1423impl ModelWorker for LoadBalancedWorker {
1424 async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError> {
1425 let idx = self.pick();
1426 self.in_flight[idx].fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1427 let _guard = InFlightGuard(&self.in_flight[idx]);
1428 let label = self
1429 .names
1430 .get(idx)
1431 .map(String::as_str)
1432 .unwrap_or("unknown");
1433 crate::metrics::set_queue_depth(
1434 label,
1435 self.in_flight[idx].load(std::sync::atomic::Ordering::Relaxed) as i64,
1436 );
1437 self.workers[idx].infer(prompt).await
1438 }
1439
1440 async fn infer_stream(&self, prompt: &str) -> Result<TokenStream, OrchestratorError> {
1441 let idx = self.pick();
1442 self.in_flight[idx].fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1443 let result = self.workers[idx].infer_stream(prompt).await;
1444 self.in_flight[idx].fetch_sub(1, std::sync::atomic::Ordering::Relaxed);
1449 result
1450 }
1451}
1452
1453#[cfg(test)]
1454mod tests {
1455 use super::*;
1456 use parking_lot::Mutex;
1457 use wiremock::matchers::{header, method, path};
1458 use wiremock::{Mock, MockServer, ResponseTemplate};
1459
1460 static ENV_MUTEX: Mutex<()> = Mutex::new(());
1462
1463 fn make_openai_worker_for(base_url: &str) -> OpenAiWorker {
1468 std::env::set_var("OPENAI_API_KEY", "test-key-openai");
1469 let w = OpenAiWorker::new("gpt-3.5-turbo-instruct")
1470 .expect("OpenAiWorker::new must succeed when OPENAI_API_KEY is set")
1471 .with_base_url(base_url);
1472 std::env::remove_var("OPENAI_API_KEY");
1473 w
1474 }
1475
1476 fn make_anthropic_worker_for(base_url: &str) -> AnthropicWorker {
1479 std::env::set_var("ANTHROPIC_API_KEY", "test-key-anthropic");
1480 let w = AnthropicWorker::new("claude-instant-1-2")
1481 .expect("AnthropicWorker::new must succeed when ANTHROPIC_API_KEY is set")
1482 .with_base_url(base_url);
1483 std::env::remove_var("ANTHROPIC_API_KEY");
1484 w
1485 }
1486
1487 fn openai_success_body() -> serde_json::Value {
1488 serde_json::json!({"choices": [{"message": {"role": "assistant", "content": "hello world response"}}]})
1489 }
1490
1491 fn anthropic_success_body() -> serde_json::Value {
1492 serde_json::json!({
1494 "id": "msg_test",
1495 "type": "message",
1496 "role": "assistant",
1497 "content": [{"type": "text", "text": "hello world response"}],
1498 "model": "claude-3-5-sonnet-20241022",
1499 "stop_reason": "end_turn",
1500 "usage": {"input_tokens": 10, "output_tokens": 3}
1501 })
1502 }
1503
1504 fn llamacpp_success_body() -> serde_json::Value {
1505 serde_json::json!({"content": "hello world response"})
1506 }
1507
1508 fn vllm_success_body() -> serde_json::Value {
1509 serde_json::json!({"text": ["hello world response"]})
1510 }
1511
1512 #[tokio::test]
1515 async fn test_echo_worker_infer_splits_on_whitespace() {
1516 let worker = EchoWorker::with_delay(0);
1517 let tokens = worker.infer("hello world").await.unwrap();
1518 assert_eq!(tokens, vec!["hello", "world"]);
1519 }
1520
1521 #[tokio::test]
1522 async fn test_echo_worker_infer_empty_prompt_returns_empty_tokens() {
1523 let worker = EchoWorker::with_delay(0);
1524 let tokens = worker.infer("").await.unwrap();
1525 assert!(tokens.is_empty(), "empty prompt should produce no tokens");
1526 }
1527
1528 #[tokio::test]
1529 async fn test_echo_worker_infer_single_word_returns_one_token() {
1530 let worker = EchoWorker::with_delay(0);
1531 let tokens = worker.infer("hello").await.unwrap();
1532 assert_eq!(tokens, vec!["hello"]);
1533 }
1534
1535 #[tokio::test]
1536 async fn test_echo_worker_infer_multiple_whitespace_is_normalised() {
1537 let worker = EchoWorker::with_delay(0);
1539 let tokens = worker.infer("a b c").await.unwrap();
1540 assert_eq!(tokens, vec!["a", "b", "c"]);
1541 }
1542
1543 #[tokio::test]
1544 async fn test_echo_worker_with_delay_stores_delay_ms() {
1545 let worker = EchoWorker::with_delay(42);
1546 assert_eq!(worker.delay_ms, 42);
1547 }
1548
1549 #[tokio::test]
1550 async fn test_echo_worker_new_delay_is_10ms() {
1551 let worker = EchoWorker::new();
1552 assert_eq!(worker.delay_ms, 10);
1553 }
1554
1555 #[tokio::test]
1556 async fn test_echo_worker_default_via_trait_works() {
1557 let worker = EchoWorker::default();
1558 let tokens = worker.infer("one two three").await.unwrap();
1559 assert_eq!(tokens.len(), 3);
1560 }
1561
1562 #[tokio::test]
1563 async fn test_echo_worker_infer_always_returns_ok() {
1564 let worker = EchoWorker::with_delay(0);
1565 assert!(worker.infer("anything").await.is_ok());
1567 }
1568
1569 #[test]
1572 fn test_openai_worker_new_missing_key_returns_config_error() {
1573 let _guard = ENV_MUTEX.lock();
1574 std::env::remove_var("OPENAI_API_KEY");
1575 let result = OpenAiWorker::new("gpt-4");
1576 assert!(
1577 result.is_err(),
1578 "Expected Err when OPENAI_API_KEY is not set"
1579 );
1580 match result.unwrap_err() {
1581 OrchestratorError::ConfigError(msg) => {
1582 assert!(
1583 msg.contains("OPENAI_API_KEY"),
1584 "Error should name the missing var"
1585 );
1586 }
1587 other => unreachable!("Expected ConfigError, got {:?}", other),
1588 }
1589 }
1590
1591 #[test]
1592 fn test_openai_worker_new_with_key_succeeds() {
1593 let _guard = ENV_MUTEX.lock();
1594 std::env::set_var("OPENAI_API_KEY", "sk-test");
1595 let result = OpenAiWorker::new("gpt-4");
1596 std::env::remove_var("OPENAI_API_KEY");
1597 assert!(result.is_ok(), "Expected Ok when OPENAI_API_KEY is set");
1598 }
1599
1600 #[tokio::test]
1603 async fn test_openai_infer_success_parses_response_correctly() {
1604 let server = MockServer::start().await;
1605 Mock::given(method("POST"))
1606 .and(path("/chat/completions"))
1607 .respond_with(ResponseTemplate::new(200).set_body_json(openai_success_body()))
1608 .mount(&server)
1609 .await;
1610
1611 let worker = {
1612 let _g = ENV_MUTEX.lock();
1613 make_openai_worker_for(&server.uri())
1614 };
1615 let tokens = worker.infer("test prompt").await.unwrap();
1616 assert_eq!(tokens, vec!["hello world response"]);
1617 }
1618
1619 #[tokio::test]
1620 async fn test_openai_infer_http_500_returns_inference_error() {
1621 let server = MockServer::start().await;
1622 Mock::given(method("POST"))
1623 .and(path("/chat/completions"))
1624 .respond_with(ResponseTemplate::new(500).set_body_string("internal error"))
1625 .mount(&server)
1626 .await;
1627
1628 let worker = {
1629 let _g = ENV_MUTEX.lock();
1630 make_openai_worker_for(&server.uri())
1631 };
1632 let result = worker.infer("test").await;
1633 assert!(result.is_err());
1634 match result.unwrap_err() {
1635 OrchestratorError::Inference(msg) => {
1636 assert!(
1637 msg.contains("500"),
1638 "Error message should include the status code"
1639 );
1640 }
1641 other => unreachable!("Expected Inference error, got {:?}", other),
1642 }
1643 }
1644
1645 #[tokio::test]
1646 async fn test_openai_infer_empty_choices_returns_inference_error() {
1647 let server = MockServer::start().await;
1648 Mock::given(method("POST"))
1649 .and(path("/chat/completions"))
1650 .respond_with(
1651 ResponseTemplate::new(200).set_body_json(serde_json::json!({"choices": []})),
1652 )
1653 .mount(&server)
1654 .await;
1655
1656 let worker = {
1657 let _g = ENV_MUTEX.lock();
1658 make_openai_worker_for(&server.uri())
1659 };
1660 let result = worker.infer("test").await;
1661 assert!(result.is_err());
1662 match result.unwrap_err() {
1663 OrchestratorError::Inference(msg) => {
1664 assert!(
1665 msg.contains("choices"),
1666 "Error should mention missing choices"
1667 );
1668 }
1669 other => unreachable!("Expected Inference error, got {:?}", other),
1670 }
1671 }
1672
1673 #[tokio::test]
1674 async fn test_openai_infer_invalid_json_returns_inference_error() {
1675 let server = MockServer::start().await;
1676 Mock::given(method("POST"))
1677 .and(path("/chat/completions"))
1678 .respond_with(ResponseTemplate::new(200).set_body_string("not valid json {{{{"))
1679 .mount(&server)
1680 .await;
1681
1682 let worker = {
1683 let _g = ENV_MUTEX.lock();
1684 make_openai_worker_for(&server.uri())
1685 };
1686 assert!(worker.infer("test").await.is_err());
1687 }
1688
1689 #[tokio::test]
1690 async fn test_openai_infer_sends_authorization_header() {
1691 let server = MockServer::start().await;
1692 Mock::given(method("POST"))
1696 .and(path("/chat/completions"))
1697 .and(header("authorization", "Bearer test-key-openai"))
1698 .respond_with(ResponseTemplate::new(200).set_body_json(openai_success_body()))
1699 .mount(&server)
1700 .await;
1701
1702 let worker = {
1703 let _g = ENV_MUTEX.lock();
1704 make_openai_worker_for(&server.uri())
1705 };
1706 let result = worker.infer("test").await;
1707 assert!(
1708 result.is_ok(),
1709 "Request with correct auth header should succeed"
1710 );
1711 }
1712
1713 #[tokio::test]
1714 async fn test_openai_infer_sends_correct_model_in_request_body() {
1715 let server = MockServer::start().await;
1716 Mock::given(method("POST"))
1717 .and(path("/chat/completions"))
1718 .respond_with(ResponseTemplate::new(200).set_body_json(openai_success_body()))
1719 .mount(&server)
1720 .await;
1721
1722 let worker = {
1723 let _g = ENV_MUTEX.lock();
1724 make_openai_worker_for(&server.uri())
1725 };
1726 let _ = worker.infer("test").await;
1727
1728 let reqs = server.received_requests().await.unwrap();
1729 assert_eq!(reqs.len(), 1, "Exactly one request should be sent");
1730 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
1731 assert_eq!(body["model"], "gpt-3.5-turbo-instruct");
1732 }
1733
1734 #[tokio::test]
1735 async fn test_openai_with_max_tokens_sends_correct_value() {
1736 let server = MockServer::start().await;
1737 Mock::given(method("POST"))
1738 .and(path("/chat/completions"))
1739 .respond_with(ResponseTemplate::new(200).set_body_json(openai_success_body()))
1740 .mount(&server)
1741 .await;
1742
1743 let worker = {
1744 let _g = ENV_MUTEX.lock();
1745 std::env::set_var("OPENAI_API_KEY", "test-key-openai");
1746 let w = OpenAiWorker::new("gpt-4")
1747 .unwrap()
1748 .with_max_tokens(1024)
1749 .with_base_url(&server.uri());
1750 std::env::remove_var("OPENAI_API_KEY");
1751 w
1752 };
1753 let _ = worker.infer("test").await;
1754
1755 let reqs = server.received_requests().await.unwrap();
1756 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
1757 assert_eq!(body["max_tokens"], 1024);
1758 }
1759
1760 #[tokio::test]
1761 async fn test_openai_with_temperature_sends_correct_value() {
1762 let server = MockServer::start().await;
1763 Mock::given(method("POST"))
1764 .and(path("/chat/completions"))
1765 .respond_with(ResponseTemplate::new(200).set_body_json(openai_success_body()))
1766 .mount(&server)
1767 .await;
1768
1769 let worker = {
1770 let _g = ENV_MUTEX.lock();
1771 std::env::set_var("OPENAI_API_KEY", "test-key-openai");
1772 let w = OpenAiWorker::new("gpt-4")
1773 .unwrap()
1774 .with_temperature(0.3)
1775 .with_base_url(&server.uri());
1776 std::env::remove_var("OPENAI_API_KEY");
1777 w
1778 };
1779 let _ = worker.infer("test").await;
1780
1781 let reqs = server.received_requests().await.unwrap();
1782 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
1783 let temp = body["temperature"].as_f64().unwrap();
1784 assert!(
1785 (temp - 0.3_f64).abs() < 0.01,
1786 "Temperature should be ~0.3, got {temp}"
1787 );
1788 }
1789
1790 #[test]
1793 fn test_anthropic_worker_new_missing_key_returns_config_error() {
1794 let _guard = ENV_MUTEX.lock();
1795 std::env::remove_var("ANTHROPIC_API_KEY");
1796 let result = AnthropicWorker::new("claude-3-5-sonnet-20241022");
1797 assert!(
1798 result.is_err(),
1799 "Expected Err when ANTHROPIC_API_KEY is not set"
1800 );
1801 match result.unwrap_err() {
1802 OrchestratorError::ConfigError(msg) => {
1803 assert!(
1804 msg.contains("ANTHROPIC_API_KEY"),
1805 "Error should name the missing var"
1806 );
1807 }
1808 other => unreachable!("Expected ConfigError, got {:?}", other),
1809 }
1810 }
1811
1812 #[test]
1813 fn test_anthropic_worker_new_with_key_succeeds() {
1814 let _guard = ENV_MUTEX.lock();
1815 std::env::set_var("ANTHROPIC_API_KEY", "sk-ant-test");
1816 let result = AnthropicWorker::new("claude-3-5-sonnet-20241022");
1817 std::env::remove_var("ANTHROPIC_API_KEY");
1818 assert!(result.is_ok(), "Expected Ok when ANTHROPIC_API_KEY is set");
1819 }
1820
1821 #[tokio::test]
1824 async fn test_anthropic_infer_success_returns_tokens() {
1825 let server = MockServer::start().await;
1826 Mock::given(method("POST"))
1827 .and(path("/messages"))
1828 .respond_with(ResponseTemplate::new(200).set_body_json(anthropic_success_body()))
1829 .mount(&server)
1830 .await;
1831
1832 let worker = {
1833 let _g = ENV_MUTEX.lock();
1834 make_anthropic_worker_for(&server.uri())
1835 };
1836 let tokens = worker.infer("test prompt").await.unwrap();
1837 assert_eq!(tokens, vec!["hello world response"]);
1838 }
1839
1840 #[tokio::test]
1841 async fn test_anthropic_infer_http_500_returns_inference_error() {
1842 let server = MockServer::start().await;
1843 Mock::given(method("POST"))
1844 .and(path("/messages"))
1845 .respond_with(ResponseTemplate::new(500).set_body_string("error"))
1846 .mount(&server)
1847 .await;
1848
1849 let worker = {
1850 let _g = ENV_MUTEX.lock();
1851 make_anthropic_worker_for(&server.uri())
1852 };
1853 let result = worker.infer("test").await;
1854 assert!(result.is_err());
1855 match result.unwrap_err() {
1856 OrchestratorError::Inference(msg) => {
1857 assert!(msg.contains("500"), "Error should include the status code");
1858 }
1859 other => unreachable!("Expected Inference error, got {:?}", other),
1860 }
1861 }
1862
1863 #[tokio::test]
1864 async fn test_anthropic_infer_invalid_json_returns_inference_error() {
1865 let server = MockServer::start().await;
1866 Mock::given(method("POST"))
1867 .and(path("/messages"))
1868 .respond_with(ResponseTemplate::new(200).set_body_string("not json"))
1869 .mount(&server)
1870 .await;
1871
1872 let worker = {
1873 let _g = ENV_MUTEX.lock();
1874 make_anthropic_worker_for(&server.uri())
1875 };
1876 assert!(worker.infer("test").await.is_err());
1877 }
1878
1879 #[tokio::test]
1880 async fn test_anthropic_infer_sends_api_key_header() {
1881 let server = MockServer::start().await;
1882 Mock::given(method("POST"))
1883 .and(path("/messages"))
1884 .and(header("x-api-key", "test-key-anthropic"))
1885 .respond_with(ResponseTemplate::new(200).set_body_json(anthropic_success_body()))
1886 .mount(&server)
1887 .await;
1888
1889 let worker = {
1890 let _g = ENV_MUTEX.lock();
1891 make_anthropic_worker_for(&server.uri())
1892 };
1893 let result = worker.infer("test").await;
1894 assert!(
1895 result.is_ok(),
1896 "Request with correct x-api-key header should succeed"
1897 );
1898 }
1899
1900 #[tokio::test]
1901 async fn test_anthropic_infer_sends_version_header() {
1902 let server = MockServer::start().await;
1903 Mock::given(method("POST"))
1904 .and(path("/messages"))
1905 .and(header("anthropic-version", "2023-06-01"))
1906 .respond_with(ResponseTemplate::new(200).set_body_json(anthropic_success_body()))
1907 .mount(&server)
1908 .await;
1909
1910 let worker = {
1911 let _g = ENV_MUTEX.lock();
1912 make_anthropic_worker_for(&server.uri())
1913 };
1914 let result = worker.infer("test").await;
1915 assert!(
1916 result.is_ok(),
1917 "Request with correct anthropic-version header should succeed"
1918 );
1919 }
1920
1921 #[tokio::test]
1922 async fn test_anthropic_infer_sends_correct_model_in_request_body() {
1923 let server = MockServer::start().await;
1924 Mock::given(method("POST"))
1925 .and(path("/messages"))
1926 .respond_with(ResponseTemplate::new(200).set_body_json(anthropic_success_body()))
1927 .mount(&server)
1928 .await;
1929
1930 let worker = {
1931 let _g = ENV_MUTEX.lock();
1932 make_anthropic_worker_for(&server.uri())
1933 };
1934 let _ = worker.infer("test").await;
1935
1936 let reqs = server.received_requests().await.unwrap();
1937 assert_eq!(reqs.len(), 1, "Exactly one request should be sent");
1938 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
1939 assert_eq!(body["model"], "claude-instant-1-2");
1940 }
1941
1942 #[tokio::test]
1943 async fn test_anthropic_with_max_tokens_sends_correct_value() {
1944 let server = MockServer::start().await;
1945 Mock::given(method("POST"))
1946 .and(path("/messages"))
1947 .respond_with(ResponseTemplate::new(200).set_body_json(anthropic_success_body()))
1948 .mount(&server)
1949 .await;
1950
1951 let worker = {
1952 let _g = ENV_MUTEX.lock();
1953 std::env::set_var("ANTHROPIC_API_KEY", "test-key-anthropic");
1954 let w = AnthropicWorker::new("claude-instant-1-2")
1955 .unwrap()
1956 .with_max_tokens(2048)
1957 .with_base_url(&server.uri());
1958 std::env::remove_var("ANTHROPIC_API_KEY");
1959 w
1960 };
1961 let _ = worker.infer("test").await;
1962
1963 let reqs = server.received_requests().await.unwrap();
1964 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
1965 assert_eq!(body["max_tokens"], 2048);
1967 }
1968
1969 #[tokio::test]
1970 async fn test_anthropic_infer_formats_prompt_with_human_and_assistant_prefix() {
1971 let server = MockServer::start().await;
1972 Mock::given(method("POST"))
1973 .and(path("/messages"))
1974 .respond_with(ResponseTemplate::new(200).set_body_json(anthropic_success_body()))
1975 .mount(&server)
1976 .await;
1977
1978 let worker = {
1979 let _g = ENV_MUTEX.lock();
1980 make_anthropic_worker_for(&server.uri())
1981 };
1982 let _ = worker.infer("my question").await;
1983
1984 let reqs = server.received_requests().await.unwrap();
1985 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
1986 let messages = body["messages"].as_array().unwrap();
1988 assert!(!messages.is_empty(), "Messages array must not be empty");
1989 let content = messages[0]["content"].as_str().unwrap();
1990 assert!(
1991 content.contains("my question"),
1992 "Prompt should include the original input"
1993 );
1994 }
1995
1996 #[test]
1999 fn test_llamacpp_default_constructor_builds_worker() {
2000 let worker = LlamaCppWorker::new();
2002 assert!(!worker.url.is_empty(), "URL should be non-empty");
2003 }
2004
2005 #[tokio::test]
2006 async fn test_llamacpp_infer_success_returns_tokens() {
2007 let server = MockServer::start().await;
2008 Mock::given(method("POST"))
2009 .and(path("/completion"))
2010 .respond_with(ResponseTemplate::new(200).set_body_json(llamacpp_success_body()))
2011 .mount(&server)
2012 .await;
2013
2014 let worker = LlamaCppWorker::new().with_url(server.uri());
2015 let tokens = worker.infer("test prompt").await.unwrap();
2016 assert_eq!(tokens, vec!["hello world response"]);
2017 }
2018
2019 #[tokio::test]
2020 async fn test_llamacpp_infer_http_500_returns_inference_error() {
2021 let server = MockServer::start().await;
2022 Mock::given(method("POST"))
2023 .and(path("/completion"))
2024 .respond_with(ResponseTemplate::new(500).set_body_string("server error"))
2025 .mount(&server)
2026 .await;
2027
2028 let worker = LlamaCppWorker::new().with_url(server.uri());
2029 let result = worker.infer("test").await;
2030 assert!(result.is_err());
2031 match result.unwrap_err() {
2032 OrchestratorError::Inference(msg) => {
2033 assert!(msg.contains("500"), "Error should include the status code");
2034 }
2035 other => unreachable!("Expected Inference error, got {:?}", other),
2036 }
2037 }
2038
2039 #[tokio::test]
2040 async fn test_llamacpp_infer_invalid_json_returns_inference_error() {
2041 let server = MockServer::start().await;
2042 Mock::given(method("POST"))
2043 .and(path("/completion"))
2044 .respond_with(ResponseTemplate::new(200).set_body_string("not json"))
2045 .mount(&server)
2046 .await;
2047
2048 let worker = LlamaCppWorker::new().with_url(server.uri());
2049 assert!(worker.infer("test").await.is_err());
2050 }
2051
2052 #[tokio::test]
2053 async fn test_llamacpp_sends_request_to_completion_endpoint() {
2054 let server = MockServer::start().await;
2055 Mock::given(method("POST"))
2056 .and(path("/completion"))
2057 .respond_with(ResponseTemplate::new(200).set_body_json(llamacpp_success_body()))
2058 .mount(&server)
2059 .await;
2060
2061 let worker = LlamaCppWorker::new().with_url(server.uri());
2062 let _ = worker.infer("test").await;
2063
2064 let reqs = server.received_requests().await.unwrap();
2065 assert_eq!(reqs.len(), 1, "Exactly one request should be sent");
2066 assert_eq!(reqs[0].url.path(), "/completion");
2067 }
2068
2069 #[tokio::test]
2070 async fn test_llamacpp_with_max_tokens_sends_n_predict_field() {
2071 let server = MockServer::start().await;
2072 Mock::given(method("POST"))
2073 .and(path("/completion"))
2074 .respond_with(ResponseTemplate::new(200).set_body_json(llamacpp_success_body()))
2075 .mount(&server)
2076 .await;
2077
2078 let worker = LlamaCppWorker::new()
2079 .with_url(server.uri())
2080 .with_max_tokens(512);
2081 let _ = worker.infer("test").await;
2082
2083 let reqs = server.received_requests().await.unwrap();
2084 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
2085 assert_eq!(body["n_predict"], 512);
2086 }
2087
2088 #[tokio::test]
2089 async fn test_llamacpp_infer_empty_content_returns_empty_tokens() {
2090 let server = MockServer::start().await;
2091 Mock::given(method("POST"))
2092 .and(path("/completion"))
2093 .respond_with(
2094 ResponseTemplate::new(200).set_body_json(serde_json::json!({"content": ""})),
2095 )
2096 .mount(&server)
2097 .await;
2098
2099 let worker = LlamaCppWorker::new().with_url(server.uri());
2100 let tokens = worker.infer("test").await.unwrap();
2101 assert!(tokens.is_empty(), "Empty content should produce no tokens");
2102 }
2103
2104 #[tokio::test]
2105 async fn test_llamacpp_with_url_overrides_default_server() {
2106 let server = MockServer::start().await;
2107 Mock::given(method("POST"))
2108 .and(path("/completion"))
2109 .respond_with(ResponseTemplate::new(200).set_body_json(llamacpp_success_body()))
2110 .mount(&server)
2111 .await;
2112
2113 let worker = LlamaCppWorker::new().with_url(server.uri());
2115 let result = worker.infer("test").await;
2116 assert!(
2117 result.is_ok(),
2118 "Request should reach the mock server via with_url"
2119 );
2120 }
2121
2122 #[test]
2125 fn test_vllm_default_constructor_builds_worker() {
2126 let worker = VllmWorker::new();
2127 assert!(!worker.url.is_empty(), "URL should be non-empty");
2128 }
2129
2130 #[tokio::test]
2131 async fn test_vllm_infer_success_returns_tokens() {
2132 let server = MockServer::start().await;
2133 Mock::given(method("POST"))
2134 .and(path("/generate"))
2135 .respond_with(ResponseTemplate::new(200).set_body_json(vllm_success_body()))
2136 .mount(&server)
2137 .await;
2138
2139 let worker = VllmWorker::new().with_url(server.uri());
2140 let tokens = worker.infer("test prompt").await.unwrap();
2141 assert_eq!(tokens, vec!["hello world response"]);
2142 }
2143
2144 #[tokio::test]
2145 async fn test_vllm_infer_http_500_returns_inference_error() {
2146 let server = MockServer::start().await;
2147 Mock::given(method("POST"))
2148 .and(path("/generate"))
2149 .respond_with(ResponseTemplate::new(500).set_body_string("server error"))
2150 .mount(&server)
2151 .await;
2152
2153 let worker = VllmWorker::new().with_url(server.uri());
2154 let result = worker.infer("test").await;
2155 assert!(result.is_err());
2156 match result.unwrap_err() {
2157 OrchestratorError::Inference(msg) => {
2158 assert!(msg.contains("500"), "Error should include the status code");
2159 }
2160 other => unreachable!("Expected Inference error, got {:?}", other),
2161 }
2162 }
2163
2164 #[tokio::test]
2165 async fn test_vllm_infer_empty_text_array_returns_inference_error() {
2166 let server = MockServer::start().await;
2167 Mock::given(method("POST"))
2168 .and(path("/generate"))
2169 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"text": []})))
2170 .mount(&server)
2171 .await;
2172
2173 let worker = VllmWorker::new().with_url(server.uri());
2174 let result = worker.infer("test").await;
2175 assert!(result.is_err());
2176 match result.unwrap_err() {
2177 OrchestratorError::Inference(msg) => {
2178 assert!(msg.contains("Empty"), "Error should mention empty response");
2179 }
2180 other => unreachable!("Expected Inference error, got {:?}", other),
2181 }
2182 }
2183
2184 #[tokio::test]
2185 async fn test_vllm_infer_invalid_json_returns_inference_error() {
2186 let server = MockServer::start().await;
2187 Mock::given(method("POST"))
2188 .and(path("/generate"))
2189 .respond_with(ResponseTemplate::new(200).set_body_string("not json"))
2190 .mount(&server)
2191 .await;
2192
2193 let worker = VllmWorker::new().with_url(server.uri());
2194 assert!(worker.infer("test").await.is_err());
2195 }
2196
2197 #[tokio::test]
2198 async fn test_vllm_sends_request_to_generate_endpoint() {
2199 let server = MockServer::start().await;
2200 Mock::given(method("POST"))
2201 .and(path("/generate"))
2202 .respond_with(ResponseTemplate::new(200).set_body_json(vllm_success_body()))
2203 .mount(&server)
2204 .await;
2205
2206 let worker = VllmWorker::new().with_url(server.uri());
2207 let _ = worker.infer("test").await;
2208
2209 let reqs = server.received_requests().await.unwrap();
2210 assert_eq!(reqs.len(), 1, "Exactly one request should be sent");
2211 assert_eq!(reqs[0].url.path(), "/generate");
2212 }
2213
2214 #[tokio::test]
2215 async fn test_vllm_with_max_tokens_sends_correct_value() {
2216 let server = MockServer::start().await;
2217 Mock::given(method("POST"))
2218 .and(path("/generate"))
2219 .respond_with(ResponseTemplate::new(200).set_body_json(vllm_success_body()))
2220 .mount(&server)
2221 .await;
2222
2223 let worker = VllmWorker::new()
2224 .with_url(server.uri())
2225 .with_max_tokens(2048);
2226 let _ = worker.infer("test").await;
2227
2228 let reqs = server.received_requests().await.unwrap();
2229 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
2230 assert_eq!(body["max_tokens"], 2048);
2231 }
2232
2233 #[tokio::test]
2234 async fn test_vllm_with_top_p_sends_correct_value() {
2235 let server = MockServer::start().await;
2236 Mock::given(method("POST"))
2237 .and(path("/generate"))
2238 .respond_with(ResponseTemplate::new(200).set_body_json(vllm_success_body()))
2239 .mount(&server)
2240 .await;
2241
2242 let worker = VllmWorker::new().with_url(server.uri()).with_top_p(0.85);
2243 let _ = worker.infer("test").await;
2244
2245 let reqs = server.received_requests().await.unwrap();
2246 let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
2247 let top_p = body["top_p"].as_f64().unwrap();
2248 assert!(
2249 (top_p - 0.85_f64).abs() < 0.01,
2250 "top_p should be ~0.85, got {top_p}"
2251 );
2252 }
2253
2254 #[tokio::test]
2255 async fn test_vllm_with_url_overrides_default_server() {
2256 let server = MockServer::start().await;
2257 Mock::given(method("POST"))
2258 .and(path("/generate"))
2259 .respond_with(ResponseTemplate::new(200).set_body_json(vllm_success_body()))
2260 .mount(&server)
2261 .await;
2262
2263 let worker = VllmWorker::new().with_url(server.uri());
2265 let result = worker.infer("test").await;
2266 assert!(
2267 result.is_ok(),
2268 "Request should reach the mock server via with_url"
2269 );
2270 }
2271}