Skip to main content

tokio_prompt_orchestrator/
worker.rs

1//! Model worker abstraction and implementations
2//!
3//! Provides the ModelWorker trait and production-ready implementations:
4//! - EchoWorker: Testing/demo worker
5//! - OpenAiWorker: OpenAI API (GPT-4, GPT-3.5, etc.)
6//! - AnthropicWorker: Anthropic Claude API
7//! - LlamaCppWorker: Local llama.cpp server
8//! - VllmWorker: vLLM inference server
9//!
10//! ## Environment Variables
11//!
12//! - `OPENAI_API_KEY`: Required for OpenAiWorker
13//! - `ANTHROPIC_API_KEY`: Required for AnthropicWorker
14//! - `LLAMA_CPP_URL`: llama.cpp server URL (default: http://localhost:8080)
15//! - `VLLM_URL`: vLLM server URL (default: http://localhost:8000)
16//!
17//! ## Design note: workers are pipeline components
18//!
19//! Workers are designed to be used through the pipeline orchestrated in
20//! `stages.rs`, not called directly in production code.  The pipeline provides
21//! circuit breaking, backpressure, dead-letter queuing, and timeout handling
22//! around every worker call.
23//!
24//! If you call a worker directly (e.g. in tests or custom integrations), you
25//! must supply your own retry, timeout, and backoff strategy.
26//!
27//! # Note (debug builds only)
28//!
29//! In debug builds (`cfg(debug_assertions)`) consider adding assertions that
30//! verify a pipeline context is present when calling workers directly, to
31//! catch accidental direct use in integration tests.
32
33use 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
41/// Boxed streaming token iterator returned by `infer_stream`.
42pub type TokenStream =
43    Pin<Box<dyn Stream<Item = Result<String, OrchestratorError>> + Send + 'static>>;
44
45/// Parse a `Retry-After` header value (seconds integer or HTTP-date) into
46/// a `Duration`.  Returns `None` if the header is absent or unparseable.
47fn parse_retry_after(headers: &reqwest::header::HeaderMap) -> Option<Duration> {
48    let value = headers.get("retry-after")?.to_str().ok()?;
49    // Prefer integer seconds; fall back to a fixed 60 s if it's an HTTP-date.
50    let secs: u64 = value.trim().parse().unwrap_or(60);
51    Some(Duration::from_secs(secs))
52}
53
54/// Read `x-ratelimit-remaining-requests` and log a warning when low.
55fn 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/// Trait for model inference workers
72///
73/// Implementations must be thread-safe (Send + Sync) for use across tasks.
74/// The trait is object-safe to allow dynamic dispatch via `Arc<dyn ModelWorker>`.
75///
76/// # Resilience
77///
78/// Worker implementations do **not** include retry logic. Retries are handled
79/// by the pipeline's inference stage (see `stages::inference_stage`). If you
80/// call a worker directly outside of the pipeline, you are responsible for
81/// implementing appropriate retry, timeout, and backoff logic.
82#[async_trait]
83pub trait ModelWorker: Send + Sync {
84    /// Perform inference on the given prompt.
85    ///
86    /// Returns tokens as a vector of strings.
87    /// For streaming implementations, this should be the final token set.
88    ///
89    /// An empty token vector is valid but callers should treat it as an empty response.
90    async fn infer(&self, prompt: &str) -> Result<Vec<String>, OrchestratorError>;
91
92    /// Stream inference tokens as they arrive from the provider.
93    ///
94    /// The default implementation calls `infer` and yields each token in order,
95    /// so workers that don't override this still work with streaming consumers.
96    /// Override for true SSE/chunked streaming from the provider.
97    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
104/// Stream tokens from any [`ModelWorker`].
105///
106/// Spawns inference in a background task and returns a [`tokio::sync::mpsc::Receiver`].
107/// Tokens arrive one at a time; the channel closes when inference completes or errors.
108///
109/// The default channel capacity is 64. For workers that return many tokens you may
110/// want a larger buffer, but 64 is sufficient for typical LLM token-by-token delivery.
111///
112/// # Example
113/// ```no_run
114/// # use std::sync::Arc;
115/// # use tokio_prompt_orchestrator::{worker::{EchoWorker, ModelWorker}, worker::stream_worker};
116/// # tokio_test::block_on(async {
117/// let worker: Arc<dyn ModelWorker> = Arc::new(EchoWorker::new());
118/// let mut rx = stream_worker(worker, "hello world".to_string());
119/// while let Some(result) = rx.recv().await {
120///     match result {
121///         Ok(token) => print!("{token}"),
122///         Err(e) => eprintln!("error: {e}"),
123///     }
124/// }
125/// # });
126/// ```
127pub 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    // Task is intentionally fire-and-forget; errors are logged above.
133    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                        // Receiver was dropped; stop sending.
139                        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
152// ============================================================================
153// Echo Worker (Testing)
154// ============================================================================
155
156/// Dummy echo worker for testing
157///
158/// Simply splits the prompt into words and returns them as tokens.
159/// Useful for pipeline smoke tests without real model dependencies.
160pub struct EchoWorker {
161    /// Simulated inference delay
162    pub delay_ms: u64,
163}
164
165impl EchoWorker {
166    /// Create a new `EchoWorker` with a default 10 ms simulated delay.
167    ///
168    /// The echo worker requires no API keys or external services. It splits the
169    /// prompt on whitespace and returns each word as a token, making it ideal
170    /// for pipeline smoke tests and local development.
171    ///
172    /// # Errors
173    ///
174    /// This constructor never returns an error.
175    ///
176    /// # Examples
177    ///
178    /// ```no_run
179    /// use tokio_prompt_orchestrator::worker::EchoWorker;
180    /// use std::sync::Arc;
181    ///
182    /// # tokio_test::block_on(async {
183    /// let worker = Arc::new(EchoWorker::new());
184    /// // Use worker with spawn_pipeline or directly
185    /// # });
186    /// ```
187    pub fn new() -> Self {
188        Self { delay_ms: 10 }
189    }
190
191    /// Create a new `EchoWorker` with a custom simulated inference delay in milliseconds.
192    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        // Simulate inference latency
207        tokio::time::sleep(tokio::time::Duration::from_millis(self.delay_ms)).await;
208
209        // Echo back the prompt as tokens
210        let tokens: Vec<String> = prompt.split_whitespace().map(str::to_string).collect();
211
212        Ok(tokens)
213    }
214}
215
216// ============================================================================
217// OpenAI Worker
218// ============================================================================
219
220/// OpenAI API request payload
221#[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/// OpenAI chat/completions API response
236#[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/// OpenAI API worker (GPT-4, GPT-3.5-turbo-instruct, etc.)
252///
253/// Requires OPENAI_API_KEY environment variable.
254///
255/// ## Example
256///
257/// ```no_run
258/// # use tokio_prompt_orchestrator::{OpenAiWorker, OrchestratorError};
259/// # use std::sync::Arc;
260/// # fn example() -> Result<(), OrchestratorError> {
261/// let worker = Arc::new(
262///     OpenAiWorker::new("gpt-3.5-turbo-instruct")?
263///         .with_max_tokens(512)
264///         .with_temperature(0.7)
265/// );
266/// # Ok(()) }
267/// ```
268///
269/// # Resilience
270///
271/// `OpenAiWorker` does not retry internally. Retry logic is handled by the
272/// pipeline's inference stage. See [`ModelWorker`] for details.
273#[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    /// API base URL — override for OpenAI-compatible endpoints or testing.
282    base_url: String,
283}
284
285impl OpenAiWorker {
286    /// Create a new `OpenAiWorker` for the given model.
287    ///
288    /// Reads the API key from the `OPENAI_API_KEY` environment variable and
289    /// constructs an HTTP client configured for the OpenAI chat/completions
290    /// endpoint. Default settings: 256 max tokens, temperature 0.7, 30 s
291    /// timeout. All defaults are overridable via the builder methods.
292    ///
293    /// # Environment Variables
294    ///
295    /// | Variable | Required | Description |
296    /// |----------|----------|-------------|
297    /// | `OPENAI_API_KEY` | Yes | Bearer token sent with every request |
298    ///
299    /// # Errors
300    ///
301    /// Returns [`OrchestratorError::ConfigError`] if `OPENAI_API_KEY` is not
302    /// set in the environment.
303    ///
304    /// # Examples
305    ///
306    /// ```no_run
307    /// use tokio_prompt_orchestrator::OpenAiWorker;
308    /// use std::sync::Arc;
309    ///
310    /// // OPENAI_API_KEY must be set in the environment.
311    /// let worker = Arc::new(
312    ///     OpenAiWorker::new("gpt-4o")
313    ///         .expect("OPENAI_API_KEY must be set")
314    ///         .with_max_tokens(512)
315    ///         .with_temperature(0.7),
316    /// );
317    /// ```
318    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    /// Set maximum tokens to generate
335    pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
336        self.max_tokens = max_tokens;
337        self
338    }
339
340    /// Set temperature (0.0 – 2.0 for OpenAI).
341    ///
342    /// Values outside `[0.0, 2.0]` are clamped and a `WARN`-level log line is
343    /// emitted.  The OpenAI API rejects values outside this range with HTTP 400.
344    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    /// Set request timeout
358    pub fn with_timeout(mut self, timeout: Duration) -> Self {
359        self.timeout = timeout;
360        self
361    }
362
363    /// Override the API base URL.
364    ///
365    /// Useful for OpenAI-compatible endpoints (Azure OpenAI, Groq, local proxies)
366    /// and for pointing at a mock server in tests.
367    /// Default: `"https://api.openai.com/v1"`.
368    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        // Return the full content as a single token rather than splitting on
434        // whitespace (which loses punctuation context).  Filter out
435        // whitespace-only responses so callers see an empty vec rather than
436        // a single blank token.
437        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        // OpenAI SSE streaming: POST with stream=true, parse `data: {...}` chunks.
465        #[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        // Convert the raw byte stream into SSE token chunks.
524        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            // SSE lines look like: `data: {...}` or `data: [DONE]`
531            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        // Wrap the stream to record TTFT on the very first token.
557        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// ============================================================================
571// Anthropic Worker
572// ============================================================================
573
574/// Anthropic Claude API worker
575///
576/// Requires ANTHROPIC_API_KEY environment variable.
577///
578/// ## Example
579///
580/// ```no_run
581/// # use tokio_prompt_orchestrator::{AnthropicWorker, OrchestratorError};
582/// # use std::sync::Arc;
583/// # fn example() -> Result<(), OrchestratorError> {
584/// let worker = Arc::new(
585///     AnthropicWorker::new("claude-3-5-sonnet-20241022")?
586///         .with_max_tokens(1024)
587///         .with_temperature(1.0)
588/// );
589/// # Ok(()) }
590/// ```
591///
592/// # Resilience
593///
594/// `AnthropicWorker` does not retry internally. Retry logic is handled by the
595/// pipeline's inference stage. See [`ModelWorker`] for details.
596#[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    /// API base URL — override for Anthropic-compatible endpoints or testing.
605    base_url: String,
606}
607
608impl AnthropicWorker {
609    /// Create a new `AnthropicWorker` for the given model.
610    ///
611    /// Reads the API key from the `ANTHROPIC_API_KEY` environment variable and
612    /// constructs an HTTP client configured for the Anthropic Messages API.
613    /// Default settings: 1024 max tokens, temperature 1.0, 60 s timeout. All
614    /// defaults are overridable via the builder methods.
615    ///
616    /// # Environment Variables
617    ///
618    /// | Variable | Required | Description |
619    /// |----------|----------|-------------|
620    /// | `ANTHROPIC_API_KEY` | Yes | API key sent via `x-api-key` header |
621    ///
622    /// # Errors
623    ///
624    /// Returns [`OrchestratorError::ConfigError`] if `ANTHROPIC_API_KEY` is
625    /// not set in the environment.
626    ///
627    /// # Examples
628    ///
629    /// ```no_run
630    /// use tokio_prompt_orchestrator::AnthropicWorker;
631    /// use std::sync::Arc;
632    ///
633    /// // ANTHROPIC_API_KEY must be set in the environment.
634    /// let worker = Arc::new(
635    ///     AnthropicWorker::new("claude-3-5-sonnet-20241022")
636    ///         .expect("ANTHROPIC_API_KEY must be set")
637    ///         .with_max_tokens(1024)
638    ///         .with_temperature(1.0),
639    /// );
640    /// ```
641    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    /// Set maximum tokens to generate
658    pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
659        self.max_tokens = max_tokens;
660        self
661    }
662
663    /// Set temperature (0.0 – 1.0 for Anthropic).
664    ///
665    /// Values outside `[0.0, 1.0]` are clamped and a `WARN`-level log line is
666    /// emitted.  The Anthropic API rejects values outside this range with HTTP 400.
667    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    /// Set request timeout
681    pub fn with_timeout(mut self, timeout: Duration) -> Self {
682        self.timeout = timeout;
683        self
684    }
685
686    /// Override the API base URL.
687    ///
688    /// Useful for Anthropic-compatible endpoints or for pointing at a mock server
689    /// in tests. Default: `"https://api.anthropic.com/v1"`.
690    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        // Use the Messages API (same as infer_stream)
701        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        // Parse Messages API response: {"content": [{"type": "text", "text": "..."}]}
747        #[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        // Collect text from all text content blocks.  Filter out
764        // whitespace-only results so callers receive an empty vec rather than
765        // a single blank token.
766        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        // Anthropic Messages API with stream=true emits SSE events.
792        // We parse `content_block_delta` events to extract text deltas.
793        #[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                // SSE data lines: `data: {...}`
861                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        // Wrap the stream to record TTFT on the very first token.
885        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// ============================================================================
899// llama.cpp Worker
900// ============================================================================
901
902/// llama.cpp server request payload
903#[derive(Debug, Serialize)]
904struct LlamaCppRequest {
905    prompt: String,
906    n_predict: i32,
907    temperature: f32,
908    stop: Vec<String>,
909}
910
911/// llama.cpp server response
912#[derive(Debug, Deserialize)]
913struct LlamaCppResponse {
914    content: String,
915}
916
917/// llama.cpp HTTP server worker
918///
919/// Connects to a llama.cpp server instance.
920/// Server URL can be set via LLAMA_CPP_URL environment variable
921/// or defaults to http://localhost:8080
922///
923/// ## Example
924///
925/// ```no_run
926/// use tokio_prompt_orchestrator::LlamaCppWorker;
927/// use std::sync::Arc;
928///
929/// let worker = Arc::new(
930///     LlamaCppWorker::new()
931///         .with_url("http://localhost:8080")
932///         .with_max_tokens(512)
933/// );
934/// ```
935pub struct LlamaCppWorker {
936    client: reqwest::Client,
937    url: String,
938    max_tokens: i32,
939    temperature: f32,
940    timeout: Duration,
941}
942
943impl LlamaCppWorker {
944    /// Create a new `LlamaCppWorker` pointing at a llama.cpp HTTP server.
945    ///
946    /// Reads the server URL from the `LLAMA_CPP_URL` environment variable. If
947    /// the variable is not set the worker falls back to
948    /// `http://localhost:8080`. Default settings: 256 max tokens, temperature
949    /// 0.8, 30 s timeout.
950    ///
951    /// # Environment Variables
952    ///
953    /// | Variable | Required | Description |
954    /// |----------|----------|-------------|
955    /// | `LLAMA_CPP_URL` | No | llama.cpp server base URL (default: `http://localhost:8080`) |
956    ///
957    /// # Errors
958    ///
959    /// This constructor never returns an error. Network errors are deferred
960    /// until [`ModelWorker::infer`] is called.
961    ///
962    /// # Examples
963    ///
964    /// ```no_run
965    /// use tokio_prompt_orchestrator::LlamaCppWorker;
966    /// use std::sync::Arc;
967    ///
968    /// // Optionally set LLAMA_CPP_URL=http://gpu-host:8080 in the environment.
969    /// let worker = Arc::new(
970    ///     LlamaCppWorker::new()
971    ///         .with_max_tokens(512)
972    ///         .with_temperature(0.8),
973    /// );
974    /// ```
975    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    /// Set server URL
989    pub fn with_url(mut self, url: impl Into<String>) -> Self {
990        self.url = url.into();
991        self
992    }
993
994    /// Set maximum tokens to generate
995    pub fn with_max_tokens(mut self, max_tokens: i32) -> Self {
996        self.max_tokens = max_tokens;
997        self
998    }
999
1000    /// Set temperature (0.0 – 2.0).
1001    ///
1002    /// Values outside `[0.0, 2.0]` are clamped and a `WARN`-level log line is
1003    /// emitted.  Most llama.cpp-compatible servers reject values outside this range.
1004    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    /// Set request timeout
1018    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        // Filter out empty strings so callers receive a genuinely-empty vec when
1071        // the model returns no text content.
1072        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// ============================================================================
1089// vLLM Worker
1090// ============================================================================
1091
1092/// vLLM server request payload
1093#[derive(Debug, Serialize)]
1094struct VllmRequest {
1095    prompt: String,
1096    max_tokens: u32,
1097    temperature: f32,
1098    top_p: f32,
1099}
1100
1101/// vLLM server response
1102#[derive(Debug, Deserialize)]
1103struct VllmResponse {
1104    text: Vec<String>,
1105}
1106
1107/// vLLM inference server worker
1108///
1109/// Connects to a vLLM server instance.
1110/// Server URL can be set via VLLM_URL environment variable
1111/// or defaults to http://localhost:8000
1112///
1113/// ## Example
1114///
1115/// ```no_run
1116/// use tokio_prompt_orchestrator::VllmWorker;
1117/// use std::sync::Arc;
1118///
1119/// let worker = Arc::new(
1120///     VllmWorker::new()
1121///         .with_url("http://localhost:8000")
1122///         .with_max_tokens(1024)
1123/// );
1124/// ```
1125pub 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    /// Create a new `VllmWorker` pointing at a vLLM inference server.
1136    ///
1137    /// Reads the server URL from the `VLLM_URL` environment variable. If the
1138    /// variable is not set the worker falls back to `http://localhost:8000`.
1139    /// Default settings: 512 max tokens, temperature 0.7, top_p 0.95, 60 s
1140    /// timeout.
1141    ///
1142    /// # Environment Variables
1143    ///
1144    /// | Variable | Required | Description |
1145    /// |----------|----------|-------------|
1146    /// | `VLLM_URL` | No | vLLM server base URL (default: `http://localhost:8000`) |
1147    ///
1148    /// # Errors
1149    ///
1150    /// This constructor never returns an error. Network errors are deferred
1151    /// until [`ModelWorker::infer`] is called.
1152    ///
1153    /// # Examples
1154    ///
1155    /// ```no_run
1156    /// use tokio_prompt_orchestrator::VllmWorker;
1157    /// use std::sync::Arc;
1158    ///
1159    /// // Optionally set VLLM_URL=http://gpu-host:8000 in the environment.
1160    /// let worker = Arc::new(
1161    ///     VllmWorker::new()
1162    ///         .with_max_tokens(1024)
1163    ///         .with_temperature(0.5),
1164    /// );
1165    /// ```
1166    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    /// Set server URL
1180    pub fn with_url(mut self, url: impl Into<String>) -> Self {
1181        self.url = url.into();
1182        self
1183    }
1184
1185    /// Set maximum tokens to generate
1186    pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
1187        self.max_tokens = max_tokens;
1188        self
1189    }
1190
1191    /// Set temperature (0.0 – 2.0).
1192    ///
1193    /// Values outside `[0.0, 2.0]` are clamped and a `WARN`-level log line is
1194    /// emitted.  Most vLLM-compatible servers reject values outside this range.
1195    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    /// Set top_p sampling parameter
1209    pub fn with_top_p(mut self, top_p: f32) -> Self {
1210        self.top_p = top_p;
1211        self
1212    }
1213
1214    /// Set request timeout
1215    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// ============================================================================
1281// Load-Balanced Worker
1282// ============================================================================
1283
1284/// Strategy for distributing requests across a pool of workers.
1285#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1286pub enum LoadBalanceStrategy {
1287    /// Rotate through workers in order.
1288    RoundRobin,
1289    /// Pick the worker with the lowest running request count.
1290    LeastLoaded,
1291}
1292
1293/// A worker pool that distributes inference requests across multiple backends.
1294///
1295/// Wraps any number of `Arc<dyn ModelWorker>` instances and routes requests
1296/// using a configurable [`LoadBalanceStrategy`].  All strategies are
1297/// thread-safe; the pool is safe to share across Tokio tasks via `Arc`.
1298///
1299/// ## Example
1300///
1301/// ```no_run
1302/// # use tokio_prompt_orchestrator::{AnthropicWorker, OpenAiWorker, OrchestratorError};
1303/// # use tokio_prompt_orchestrator::worker::LoadBalancedWorker;
1304/// use std::sync::Arc;
1305///
1306/// # fn example() -> Result<(), OrchestratorError> {
1307/// let pool = LoadBalancedWorker::round_robin(vec![
1308///     Arc::new(AnthropicWorker::new("claude-sonnet-4-6")?) as Arc<dyn tokio_prompt_orchestrator::ModelWorker>,
1309///     Arc::new(OpenAiWorker::new("gpt-4o")?) as Arc<dyn tokio_prompt_orchestrator::ModelWorker>,
1310/// ]);
1311/// # Ok(()) }
1312/// ```
1313pub struct LoadBalancedWorker {
1314    workers: Vec<Arc<dyn ModelWorker>>,
1315    strategy: LoadBalanceStrategy,
1316    /// Round-robin cursor (strategy = RoundRobin).
1317    rr_counter: std::sync::atomic::AtomicUsize,
1318    /// Per-worker in-flight request counts (strategy = LeastLoaded).
1319    in_flight: Vec<std::sync::atomic::AtomicUsize>,
1320    /// Optional human-readable names for each worker, used in metrics labels.
1321    names: Vec<String>,
1322}
1323
1324/// RAII guard that decrements the in-flight counter when dropped.
1325struct 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    /// Create a round-robin pool.
1334    ///
1335    /// # Panics
1336    ///
1337    /// Panics if `workers` is empty.
1338    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    /// Create a round-robin pool by replicating a single worker `n` times.
1353    ///
1354    /// All pool slots share the same underlying `Arc`; this is useful for
1355    /// concurrency control rather than load distribution across different backends.
1356    ///
1357    /// # Panics
1358    ///
1359    /// Panics if `n` is zero.
1360    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    /// Create a least-loaded pool.
1367    ///
1368    /// # Panics
1369    ///
1370    /// Panics if `workers` is empty.
1371    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    /// Attach human-readable names to each worker slot for metrics labels.
1386    ///
1387    /// If `names` is shorter than the pool, unnamed workers default to `"unknown"`.
1388    pub fn with_names(mut self, names: Vec<String>) -> Self {
1389        self.names = names;
1390        self
1391    }
1392
1393    /// Return the number of workers in the pool.
1394    pub fn len(&self) -> usize {
1395        self.workers.len()
1396    }
1397
1398    /// Return true if the pool is empty (should never be true after construction).
1399    pub fn is_empty(&self) -> bool {
1400        self.workers.is_empty()
1401    }
1402
1403    /// Pick the index of the worker to use for the next request.
1404    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        // Note: we decrement after obtaining the stream handle.
1445        // The stream itself runs outside the critical section intentionally —
1446        // counting only the dispatch decision avoids holding a counter across
1447        // the full streaming duration, which would bias LeastLoaded unfairly.
1448        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    /// Serialise all tests that read/write environment variables so they don't race.
1461    static ENV_MUTEX: Mutex<()> = Mutex::new(());
1462
1463    // ── Helpers ───────────────────────────────────────────────────────────────
1464
1465    /// Create an `OpenAiWorker` that points at `base_url`.
1466    /// Must be called while `ENV_MUTEX` is held.
1467    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    /// Create an `AnthropicWorker` that points at `base_url`.
1477    /// Must be called while `ENV_MUTEX` is held.
1478    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        // Messages API format: {"content": [{"type": "text", "text": "..."}]}
1493        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    // ── EchoWorker ────────────────────────────────────────────────────────────
1513
1514    #[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        // split_whitespace collapses runs of whitespace
1538        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        // EchoWorker never returns an error
1566        assert!(worker.infer("anything").await.is_ok());
1567    }
1568
1569    // ── OpenAiWorker — constructor ────────────────────────────────────────────
1570
1571    #[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    // ── OpenAiWorker — inference ──────────────────────────────────────────────
1601
1602    #[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        // The mock only matches if the Authorization header has the right value.
1693        // An unmatched request returns 404, which makes the worker return Err,
1694        // causing the final assert to fail — which is the desired test signal.
1695        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    // ── AnthropicWorker — constructor ─────────────────────────────────────────
1791
1792    #[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    // ── AnthropicWorker — inference ───────────────────────────────────────────
1822
1823    #[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        // Messages API uses `max_tokens` (not `max_tokens_to_sample` from the legacy Complete API)
1966        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        // Messages API: prompt is in body["messages"][0]["content"]
1987        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    // ── LlamaCppWorker ────────────────────────────────────────────────────────
1997
1998    #[test]
1999    fn test_llamacpp_default_constructor_builds_worker() {
2000        // Uses unwrap_or_else — always succeeds
2001        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        // If with_url works, this request reaches our mock — not localhost:8080
2114        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    // ── VllmWorker ────────────────────────────────────────────────────────────
2123
2124    #[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        // If with_url works, the request reaches our mock — not localhost:8000
2264        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}