1use anda_core::{
16 AgentOutput, BoxError, BoxPinFut, CONTENT_TYPE_JSON, CompletionRequest, Json, ToolCall,
17};
18use arc_swap::ArcSwap;
19use futures_util::StreamExt;
20use serde::de::DeserializeOwned;
21use serde::{Deserialize, Serialize};
22use std::future::Future;
23use std::time::{Duration, Instant};
24use std::{
25 collections::{BTreeSet, HashMap, hash_map::Entry},
26 error::Error,
27 fmt,
28 sync::Arc,
29};
30
31pub mod anthropic;
32pub(crate) mod driver;
33pub mod gemini;
34pub mod openai;
35pub(crate) mod raw;
36#[cfg(test)]
37pub(crate) mod test_support;
38pub mod testing;
39
40pub(crate) fn null_default<'de, D, T>(deserializer: D) -> Result<T, D::Error>
43where
44 D: serde::Deserializer<'de>,
45 T: Default + Deserialize<'de>,
46{
47 Ok(Option::<T>::deserialize(deserializer)?.unwrap_or_default())
48}
49
50macro_rules! string_enum_serde {
55 ($ty:ident, { $($wire:literal $(| $alias:literal)* => $variant:ident),+ $(,)? }, $unknown:ident) => {
56 impl $ty {
57 fn as_str(&self) -> &str {
58 match self {
59 $(Self::$variant => $wire,)+
60 Self::$unknown(value) => value.as_str(),
61 }
62 }
63 }
64
65 impl serde::Serialize for $ty {
66 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
67 where
68 S: serde::Serializer,
69 {
70 serializer.serialize_str(self.as_str())
71 }
72 }
73
74 impl<'de> serde::Deserialize<'de> for $ty {
75 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
76 where
77 D: serde::Deserializer<'de>,
78 {
79 let value = String::deserialize(deserializer)?;
80 Ok(match value.as_str() {
81 $($wire $(| $alias)* => Self::$variant,)+
82 _ => Self::$unknown(value),
83 })
84 }
85 }
86 };
87}
88pub(crate) use string_enum_serde;
89
90pub(crate) fn resolve_endpoint(endpoint: Option<String>, default: &str) -> String {
93 match endpoint {
94 Some(endpoint) if !endpoint.is_empty() => endpoint,
95 _ => default.to_string(),
96 }
97}
98
99pub use reqwest;
100pub use reqwest::Proxy;
101
102use crate::APP_USER_AGENT;
103
104pub use anda_core::ModelEffort;
105
106const MODEL_REQUEST_MAX_RETRIES: usize = 3;
107const MODEL_RETRY_BACKOFF: Duration = Duration::from_secs(1);
108const MODEL_RETRY_MAX_BACKOFF: Duration = Duration::from_secs(300);
109const COMPLETION_HTTP2_KEEP_ALIVE_INTERVAL: Option<Duration> = None;
110const COMPLETION_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
111const COMPLETION_READ_TIMEOUT: Duration = Duration::from_secs(180);
112const COMPLETION_REQUEST_TIMEOUT: Duration = Duration::from_secs(600);
113
114const MAX_COMPLETION_RESPONSE_BYTES: usize = 64 * 1024 * 1024;
120
121const MAX_ERROR_BODY_BYTES: usize = 8 * 1024;
126
127pub(crate) fn error_body_excerpt(data: &[u8]) -> String {
130 if data.len() <= MAX_ERROR_BODY_BYTES {
131 return String::from_utf8_lossy(data).into_owned();
132 }
133
134 let mut text = String::from_utf8_lossy(&data[..MAX_ERROR_BODY_BYTES]).into_owned();
135 text.push_str("… [truncated]");
136 text
137}
138
139async fn read_error_body(response: reqwest::Response) -> Result<String, reqwest::Error> {
142 let mut stream = response.bytes_stream();
143 let mut body = Vec::new();
144 while let Some(chunk) = stream.next().await {
145 body.extend_from_slice(&chunk?);
146 if body.len() > MAX_ERROR_BODY_BYTES {
147 break;
148 }
149 }
150 Ok(error_body_excerpt(&body))
151}
152
153#[derive(Default, Clone, Deserialize, Serialize)]
158pub struct ModelConfig {
159 pub family: String,
161
162 pub model: String,
164
165 pub api_base: String,
167
168 pub api_key: String,
170
171 #[serde(default)]
176 pub labels: Vec<String>,
177
178 #[serde(default)]
179 pub context_window: usize,
181
182 #[serde(default)]
183 pub max_output: usize,
189
190 #[serde(default)]
195 pub effort: Option<ModelEffort>,
196
197 #[serde(default)]
199 pub disabled: bool,
200
201 #[serde(default)]
204 pub bearer_auth: bool,
205
206 #[serde(default)]
207 pub stream: bool,
211}
212
213impl std::fmt::Debug for ModelConfig {
214 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
215 f.debug_struct("ModelConfig")
216 .field("family", &self.family)
217 .field("model", &self.model)
218 .field("api_base", &self.api_base)
219 .field(
220 "api_key",
221 &if self.api_key.is_empty() {
222 ""
223 } else {
224 "[REDACTED]"
225 },
226 )
227 .field("labels", &self.labels)
228 .field("context_window", &self.context_window)
229 .field("max_output", &self.max_output)
230 .field("effort", &self.effort)
231 .field("disabled", &self.disabled)
232 .field("bearer_auth", &self.bearer_auth)
233 .field("stream", &self.stream)
234 .finish()
235 }
236}
237
238impl ModelConfig {
239 pub fn model(&self, http_client: reqwest::Client) -> Result<Model, BoxError> {
241 if self.disabled {
242 return Err("model is disabled".into());
243 }
244 if self.model.is_empty() {
245 return Err(format!("{}: model name is required", self.model).into());
246 }
247 if self.family.is_empty() {
248 return Err(format!("{}: model family is required", self.model).into());
249 }
250 if self.api_base.is_empty() {
251 return Err(format!("{}: api_base is required", self.model).into());
252 }
253 if self.api_key.is_empty() {
254 return Err(format!("{}: api_key is required", self.model).into());
255 }
256
257 let mut model = match self.family.as_str() {
258 "gemini" => Model::with_completer(Arc::new(
259 gemini::Client::new_with_client(
260 &self.api_key,
261 Some(self.api_base.clone()),
262 http_client,
263 )
264 .completion_model(&self.model)
265 .with_stream(self.stream)
266 .with_effort(self.effort)
267 .with_max_output(self.max_output),
268 )),
269 "anthropic" => {
270 let mut cli = anthropic::Client::new_with_client(
271 &self.api_key,
272 Some(self.api_base.clone()),
273 http_client,
274 );
275 if self.bearer_auth {
276 cli = cli.with_bearer_auth(true);
277 }
278 Model::with_completer(Arc::new(
279 cli.completion_model(&self.model)
280 .with_stream(self.stream)
281 .with_effort(self.effort)
282 .with_max_output(self.max_output),
283 ))
284 }
285 "openai-response" => {
286 let cli = openai::Client::new_with_client(
287 &self.api_key,
288 Some(self.api_base.clone()),
289 http_client,
290 );
291 Model::with_completer(Arc::new(
292 cli.completion_model_v2(&self.model)
293 .with_stream(self.stream)
294 .with_effort(self.effort),
295 ))
296 }
297 "openai" => {
298 let cli = openai::Client::new_with_client(
299 &self.api_key,
300 Some(self.api_base.clone()),
301 http_client,
302 );
303 if self.model.starts_with("gpt") {
304 Model::with_completer(Arc::new(
305 cli.completion_model_v2(&self.model)
306 .with_stream(self.stream)
307 .with_effort(self.effort),
308 ))
309 } else {
310 Model::with_completer(Arc::new(
311 cli.completion_model(&self.model)
312 .with_stream(self.stream)
313 .with_effort(self.effort),
314 ))
315 }
316 }
317 _ => return Err(format!("unsupported model family: {}", self.family).into()),
318 };
319
320 let labels = if self.labels.is_empty() {
321 vec![self.model.to_ascii_lowercase()]
322 } else {
323 self.labels.clone()
324 };
325 model.context_window = self.context_window;
326 model.max_output = self.max_output;
327 Ok(model.with_labels(labels))
328 }
329}
330
331pub struct Models {
341 model: ArcSwap<Option<Model>>,
342 models: ArcSwap<HashMap<String, Vec<Model>>>,
343}
344
345impl Default for Models {
346 fn default() -> Self {
347 Self {
348 model: ArcSwap::new(Arc::new(None)),
349 models: ArcSwap::new(Arc::new(HashMap::new())),
350 }
351 }
352}
353
354impl Models {
355 pub fn from_clone(other: &Models) -> Self {
357 Self {
358 model: ArcSwap::new(other.model.load_full()),
359 models: ArcSwap::new(Arc::new(other.models.load().as_ref().clone())),
360 }
361 }
362
363 pub fn replace(&self, other: &Models) {
365 self.model.store(other.model.load_full());
366 self.models
367 .store(Arc::new(other.models.load().as_ref().clone()));
368 }
369
370 pub fn from_configs(configs: &[ModelConfig], http_client: reqwest::Client) -> Self {
376 let models = Self::default();
377 for config in configs {
378 if config.disabled {
379 log::info!(
380 "skipping disabled model: family={}, model={}",
381 config.family,
382 config.model
383 );
384 continue;
385 }
386 match config.model(http_client.clone()) {
387 Ok(model) => models.inner_set(model.labels.clone(), model),
388 Err(err) => {
389 log::warn!(
390 "skipping misconfigured model: family={}, model={}, error={}",
391 config.family,
392 config.model,
393 err
394 );
395 }
396 }
397 }
398 models
399 }
400
401 pub fn contains(&self, label: &str) -> bool {
403 self.models.load().contains_key(&label.to_ascii_lowercase())
404 }
405
406 pub fn model_names(&self) -> BTreeSet<String> {
408 self.models
409 .load()
410 .values()
411 .flatten()
412 .map(|m| m.model_name())
413 .collect()
414 }
415
416 pub fn set_model(&self, model: Model) {
418 self.inner_set(model.labels.clone(), model.clone());
419 self.model.store(Arc::new(Some(model)));
420 }
421
422 pub fn set(&self, label: String, model: Model) {
428 self.inner_set(vec![label], model);
429 }
430
431 fn inner_set(&self, mut labels: Vec<String>, model: Model) {
432 if self.model.load().is_none() {
433 self.model.store(Arc::new(Some(model.clone())));
434 }
435
436 let model_name = model.model_name();
437 labels.push(model_name.to_ascii_lowercase());
438 for label in labels.iter_mut() {
439 label.make_ascii_lowercase();
440 if label == "primary" {
441 self.model.store(Arc::new(Some(model.clone())));
442 }
443 }
444
445 self.models.rcu(|models| {
447 let mut models = models.as_ref().clone();
448 for label in &labels {
449 match models.entry(label.clone()) {
450 Entry::Vacant(e) => {
451 e.insert(vec![model.clone()]);
452 }
453 Entry::Occupied(mut e) => {
454 e.get_mut().retain(|m| m.model_name() != model_name);
455 e.get_mut().push(model.clone());
456 }
457 }
458 }
459 models
460 });
461 }
462
463 pub fn get(&self, label: &str) -> Option<Model> {
467 self.models
468 .load()
469 .get(&label.to_ascii_lowercase())
470 .and_then(|v| v.last().cloned())
471 }
472
473 pub fn get_model(&self) -> Option<Model> {
476 if let Some(m) = self.model.load().as_ref() {
477 return Some(m.clone());
478 }
479 self.models
480 .load()
481 .values()
482 .next()
483 .and_then(|v| v.last().cloned())
484 }
485
486 pub fn resolve(&self, label: &str) -> Option<Model> {
492 if label.is_empty() {
493 return self.get_model();
494 }
495 self.get(label).or_else(|| self.get_model())
496 }
497}
498
499pub trait CompletionFeaturesDyn: Send + Sync + 'static {
501 fn completion(&self, req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>>;
507
508 fn model_name(&self) -> String;
510
511 fn prune_unanswered_tool_calls(&self, raw_history: &mut Vec<Json>, start: usize) {
523 raw::prune_unanswered_tool_calls(raw_history, start);
524 }
525
526 fn prune_tool_interactions(&self, raw_history: &mut Vec<Json>) {
538 raw::prune_tool_interactions(raw_history);
539 }
540}
541
542#[derive(Clone, Debug)]
544pub struct NotImplemented;
545
546impl CompletionFeaturesDyn for NotImplemented {
547 fn model_name(&self) -> String {
548 "not_implemented".to_string()
549 }
550
551 fn completion(&self, _req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>> {
552 Box::pin(futures::future::ready(Err("not implemented".into())))
553 }
554}
555
556#[derive(Clone, Debug)]
558pub struct MockImplemented;
559
560impl CompletionFeaturesDyn for MockImplemented {
561 fn model_name(&self) -> String {
562 "mock_implemented".to_string()
563 }
564
565 fn completion(&self, req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>> {
566 Box::pin(futures::future::ready(Ok(AgentOutput {
567 content: req.prompt.clone(),
568 tool_calls: req
569 .tools
570 .iter()
571 .filter_map(|tool| {
572 if req.prompt.is_empty() {
573 return None;
574 }
575 Some(ToolCall {
576 name: tool.name.clone(),
577 args: serde_json::from_str(&req.prompt).unwrap_or_default(),
578 call_id: None,
579 result: None,
580 remote_id: None,
581 })
582 })
583 .collect(),
584 ..Default::default()
585 })))
586 }
587}
588
589#[derive(Clone)]
591pub struct Model {
592 pub completer: Arc<dyn CompletionFeaturesDyn>,
594
595 pub labels: Vec<String>,
597
598 pub context_window: usize,
600
601 pub max_output: usize,
604}
605
606impl Model {
607 pub fn new(completer: Arc<dyn CompletionFeaturesDyn>) -> Self {
609 Self {
610 completer,
611 labels: Vec::new(),
612 context_window: 0,
613 max_output: 0,
614 }
615 }
616
617 pub fn with_completer(completer: Arc<dyn CompletionFeaturesDyn>) -> Self {
619 Self::new(completer)
620 }
621
622 pub fn with_labels(mut self, labels: Vec<String>) -> Self {
624 self.labels = labels;
625 self
626 }
627
628 pub fn not_implemented() -> Self {
630 Self::new(Arc::new(NotImplemented))
631 }
632
633 pub fn mock_implemented() -> Self {
635 Self::new(Arc::new(MockImplemented))
636 }
637
638 pub fn model_name(&self) -> String {
640 self.completer.model_name()
641 }
642
643 pub async fn completion(&self, mut req: CompletionRequest) -> Result<AgentOutput, BoxError> {
650 if self.max_output > 0
651 && let Some(tokens) = req.max_output_tokens.as_mut()
652 {
653 *tokens = (*tokens).min(self.max_output);
654 }
655 self.completer.completion(req).await
656 }
657
658 pub fn prune_unanswered_tool_calls(&self, raw_history: &mut Vec<Json>, start: usize) {
663 self.completer
664 .prune_unanswered_tool_calls(raw_history, start);
665 }
666
667 pub fn prune_tool_interactions(&self, raw_history: &mut Vec<Json>) {
672 self.completer.prune_tool_interactions(raw_history);
673 }
674}
675
676#[derive(Debug)]
679pub struct ModelError {
680 message: String,
681 retryable: bool,
682 status: Option<http::StatusCode>,
683 retry_after: Option<Duration>,
684 source: Option<BoxError>,
685}
686
687impl ModelError {
688 pub fn new(message: impl Into<String>) -> Self {
690 Self {
691 message: message.into(),
692 retryable: false,
693 status: None,
694 retry_after: None,
695 source: None,
696 }
697 }
698
699 pub fn with_retryable(mut self, retryable: bool) -> Self {
701 self.retryable = retryable;
702 self
703 }
704
705 pub fn with_status(mut self, status: http::StatusCode) -> Self {
707 self.status = Some(status);
708 self
709 }
710
711 pub fn with_retry_after(mut self, retry_after: Option<Duration>) -> Self {
713 self.retry_after = retry_after;
714 self
715 }
716
717 pub fn with_source(mut self, source: BoxError) -> Self {
719 self.source = Some(source);
720 self
721 }
722
723 pub fn is_retryable(&self) -> bool {
725 self.retryable
726 }
727
728 pub fn status(&self) -> Option<http::StatusCode> {
730 self.status
731 }
732
733 pub fn retry_after(&self) -> Option<Duration> {
735 self.retry_after
736 }
737}
738
739impl fmt::Display for ModelError {
740 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
741 f.write_str(&self.message)
742 }
743}
744
745impl Error for ModelError {
746 fn source(&self) -> Option<&(dyn Error + 'static)> {
747 self.source
748 .as_deref()
749 .map(|source| source as &(dyn Error + 'static))
750 }
751}
752
753fn find_in_error_chain<T>(
755 error: &(dyn Error + 'static),
756 f: impl Fn(&(dyn Error + 'static)) -> Option<T>,
757) -> Option<T> {
758 let mut current = Some(error);
759 while let Some(error) = current {
760 if let Some(found) = f(error) {
761 return Some(found);
762 }
763 current = error.source();
764 }
765 None
766}
767
768pub fn is_retryable_model_error(error: &(dyn Error + 'static)) -> bool {
777 find_in_error_chain(error, |error| {
778 let retryable = error
779 .downcast_ref::<ModelError>()
780 .is_some_and(ModelError::is_retryable)
781 || error
782 .downcast_ref::<reqwest::Error>()
783 .is_some_and(is_retryable_reqwest_error);
784 retryable.then_some(())
785 })
786 .is_some()
787}
788
789pub fn is_retryable_box_error(error: &BoxError) -> bool {
791 is_retryable_model_error(error.as_ref() as &(dyn Error + 'static))
792}
793
794pub fn model_error_status(error: &(dyn Error + 'static)) -> Option<http::StatusCode> {
796 find_in_error_chain(error, |error| error.downcast_ref::<ModelError>()?.status())
797}
798
799pub fn model_error_retry_after(error: &(dyn Error + 'static)) -> Option<Duration> {
801 find_in_error_chain(error, |error| {
802 error.downcast_ref::<ModelError>()?.retry_after()
803 })
804}
805
806pub fn is_retryable_status(status: http::StatusCode) -> bool {
809 matches!(
810 status,
811 http::StatusCode::REQUEST_TIMEOUT
812 | http::StatusCode::TOO_MANY_REQUESTS
813 | http::StatusCode::INTERNAL_SERVER_ERROR
814 | http::StatusCode::BAD_GATEWAY
815 | http::StatusCode::SERVICE_UNAVAILABLE
816 | http::StatusCode::GATEWAY_TIMEOUT
817 ) || status.as_u16() == 529
818}
819
820pub(crate) fn is_retryable_reqwest_error(err: &reqwest::Error) -> bool {
821 err.is_timeout()
822 || err.is_connect()
823 || err.is_body()
824 || err.is_decode()
825 || err.status().is_some_and(is_retryable_status)
826}
827
828pub(crate) fn format_error_chain(err: &(dyn Error + 'static)) -> String {
834 let mut message = err.to_string();
835 let mut source = err.source();
836 while let Some(err) = source {
837 let text = err.to_string();
838 if !message.contains(&text) {
840 message.push_str(": ");
841 message.push_str(&text);
842 }
843 source = err.source();
844 }
845 message
846}
847
848fn upstream_request_id(headers: &http::HeaderMap) -> Option<String> {
852 ["x-request-id", "request-id", "x-amzn-requestid", "cf-ray"]
853 .into_iter()
854 .find_map(|name| headers.get(name)?.to_str().ok())
855 .map(str::to_string)
856}
857
858pub(crate) fn completion_transport_error(
859 model: &str,
860 action: &str,
861 err: reqwest::Error,
862) -> BoxError {
863 let retryable = is_retryable_reqwest_error(&err);
864 let message = format!(
865 "{action}, model: {model}, error: {}",
866 format_error_chain(&err)
867 );
868 Box::new(
869 ModelError::new(message)
870 .with_retryable(retryable)
871 .with_source(Box::new(err)),
872 )
873}
874
875pub(crate) async fn read_completion_response_bytes(
876 response: reqwest::Response,
877 model: &str,
878) -> Result<bytes::Bytes, BoxError> {
879 let request_id = upstream_request_id(response.headers());
880 if let Some(len) = response.content_length()
882 && len > MAX_COMPLETION_RESPONSE_BYTES as u64
883 {
884 return Err(completion_response_too_large(
885 model,
886 request_id.as_deref(),
887 len as usize,
888 ));
889 }
890
891 let mut stream = response.bytes_stream();
894 let mut body: Vec<u8> = Vec::new();
895 while let Some(chunk) = stream.next().await {
896 let chunk = chunk.map_err(|err| {
897 let action = format!(
898 "Failed to read completion response (request_id: {})",
899 request_id.as_deref().unwrap_or("-")
900 );
901 completion_transport_error(model, &action, err)
902 })?;
903 if body.len() + chunk.len() > MAX_COMPLETION_RESPONSE_BYTES {
904 return Err(completion_response_too_large(
905 model,
906 request_id.as_deref(),
907 body.len() + chunk.len(),
908 ));
909 }
910 body.extend_from_slice(&chunk);
911 }
912 Ok(body.into())
913}
914
915fn completion_response_too_large(
917 model: &str,
918 request_id: Option<&str>,
919 received: usize,
920) -> BoxError {
921 Box::new(ModelError::new(format!(
922 "Completion response too large (model: {}, request_id: {}, received: {} bytes, limit: {} bytes)",
923 model,
924 request_id.unwrap_or("-"),
925 received,
926 MAX_COMPLETION_RESPONSE_BYTES,
927 )))
928}
929
930pub(crate) async fn execute_completion_request_with_retry<T, BuildRequest, HandleResponse, Fut>(
931 model: &str,
932 build_request: BuildRequest,
933 handle_response: HandleResponse,
934) -> Result<T, BoxError>
935where
936 BuildRequest: Fn() -> reqwest::RequestBuilder,
937 HandleResponse: Fn(reqwest::Response) -> Fut,
938 Fut: Future<Output = Result<T, BoxError>>,
939{
940 for attempt in 0..=MODEL_REQUEST_MAX_RETRIES {
941 let response = match build_request().send().await {
942 Ok(response) => response,
943 Err(err) => {
944 let retryable = is_retryable_reqwest_error(&err);
945 let message = format!(
946 "Failed to send completion request, model: {}, error: {}",
947 model,
948 format_error_chain(&err)
949 );
950 if retryable && attempt < MODEL_REQUEST_MAX_RETRIES {
951 log_completion_retry(model, attempt + 1);
952 backoff_before_retry(None).await;
953 continue;
954 }
955
956 return Err(Box::new(
957 ModelError::new(message)
958 .with_retryable(retryable)
959 .with_source(Box::new(err)),
960 ));
961 }
962 };
963
964 let status = response.status();
965 if status.is_success() {
966 match handle_response(response).await {
967 Ok(output) => return Ok(output),
968 Err(err) if is_retryable_box_error(&err) && attempt < MODEL_REQUEST_MAX_RETRIES => {
969 log_completion_retry(model, attempt + 1);
970 backoff_before_retry(None).await;
971 continue;
972 }
973 Err(err) => return Err(err),
974 }
975 }
976
977 let retryable = is_retryable_status(status);
978 let retry_after = retry_after_duration(response.headers());
979 let body = match read_error_body(response).await {
980 Ok(body) => body,
981 Err(err) => {
982 let retryable = retryable || is_retryable_reqwest_error(&err);
983 let message = format!(
984 "Completion failed, model: {}, status: {}; failed to read error body: {}",
985 model,
986 status,
987 format_error_chain(&err)
988 );
989 if retryable && attempt < MODEL_REQUEST_MAX_RETRIES {
990 log_completion_retry(model, attempt + 1);
991 backoff_before_retry(retry_after).await;
992 continue;
993 }
994
995 return Err(Box::new(
996 ModelError::new(message)
997 .with_retryable(retryable)
998 .with_status(status)
999 .with_retry_after(retry_after)
1000 .with_source(Box::new(err)),
1001 ));
1002 }
1003 };
1004 let message = format!(
1005 "Completion failed, model: {}, status: {}, body: {}",
1006 model, status, body
1007 );
1008
1009 if retryable && attempt < MODEL_REQUEST_MAX_RETRIES {
1010 log_completion_retry(model, attempt + 1);
1011 backoff_before_retry(retry_after).await;
1012 continue;
1013 }
1014
1015 return Err(Box::new(
1016 ModelError::new(message)
1017 .with_retryable(retryable)
1018 .with_status(status)
1019 .with_retry_after(retry_after),
1020 ));
1021 }
1022
1023 unreachable!("completion retry loop always returns before exhausting attempts")
1024}
1025
1026async fn backoff_before_retry(retry_after: Option<Duration>) {
1031 let delay = retry_after
1032 .unwrap_or(MODEL_RETRY_BACKOFF)
1033 .min(MODEL_RETRY_MAX_BACKOFF);
1034 tokio::time::sleep(delay).await;
1035}
1036
1037fn retry_after_duration(headers: &http::HeaderMap) -> Option<Duration> {
1038 let value = headers
1039 .get(http::header::RETRY_AFTER)?
1040 .to_str()
1041 .ok()?
1042 .trim();
1043 if let Ok(seconds) = value.parse::<u64>() {
1044 return Some(Duration::from_secs(seconds));
1045 }
1046
1047 let when = chrono::DateTime::parse_from_rfc2822(value).ok()?;
1050 (when.with_timezone(&chrono::Utc) - chrono::Utc::now())
1051 .to_std()
1052 .ok()
1053}
1054
1055fn log_completion_retry(model: &str, retry: usize) {
1056 log::warn!(
1058 "Retrying completion request, model: {}, retry: {}/{}",
1059 model,
1060 retry,
1061 MODEL_REQUEST_MAX_RETRIES
1062 );
1063}
1064
1065#[derive(Clone, Copy, Debug)]
1067pub struct AnyHost;
1068
1069impl PartialEq<&str> for AnyHost {
1070 fn eq(&self, _other: &&str) -> bool {
1071 true
1072 }
1073}
1074
1075pub fn request_client_builder() -> reqwest::ClientBuilder {
1077 reqwest::Client::builder()
1078 .use_rustls_tls()
1079 .https_only(true)
1080 .retry(
1081 reqwest::retry::for_host(AnyHost)
1082 .max_retries_per_request(1)
1083 .classify_fn(|req_rep| {
1084 let is_idempotent = matches!(
1085 req_rep.method(),
1086 &http::Method::GET
1087 | &http::Method::HEAD
1088 | &http::Method::OPTIONS
1089 | &http::Method::TRACE
1090 | &http::Method::PUT
1091 | &http::Method::DELETE
1092 );
1093
1094 if !is_idempotent {
1095 return req_rep.success();
1096 }
1097
1098 if req_rep.error().is_some() {
1099 return req_rep.retryable();
1100 }
1101
1102 match req_rep.status() {
1103 Some(status) if is_retryable_status(status) => req_rep.retryable(),
1104 _ => req_rep.success(),
1105 }
1106 }),
1107 )
1108 .http2_keep_alive_interval(COMPLETION_HTTP2_KEEP_ALIVE_INTERVAL)
1115 .connect_timeout(COMPLETION_CONNECT_TIMEOUT)
1116 .read_timeout(COMPLETION_READ_TIMEOUT)
1121 .timeout(COMPLETION_REQUEST_TIMEOUT)
1125 .user_agent(APP_USER_AGENT)
1126 .default_headers({
1127 let mut headers = reqwest::header::HeaderMap::new();
1128 let ct: http::HeaderValue = http::HeaderValue::from_static(CONTENT_TYPE_JSON);
1129 headers.insert(http::header::CONTENT_TYPE, ct.clone());
1130 headers.insert(http::header::ACCEPT, ct);
1131 headers
1132 })
1133}
1134
1135const SSE_DONE_MARKER: &[u8] = b"data: [DONE]";
1136
1137pub(crate) async fn read_completion_stream<T>(
1138 response: reqwest::Response,
1139 model: &str,
1140) -> Result<(Vec<T>, bool), BoxError>
1141where
1142 T: DeserializeOwned,
1143{
1144 let request_id = upstream_request_id(response.headers());
1145 let started = Instant::now();
1146 let mut body = Vec::new();
1147 let mut scanned: usize = 0;
1148 let mut stream = response.bytes_stream();
1149
1150 while let Some(chunk) = stream.next().await {
1151 let chunk = chunk.map_err(|err| {
1152 let action = format!(
1157 "Failed to read streaming completion response (request_id: {}, received: {} bytes, elapsed: {:.1?})",
1158 request_id.as_deref().unwrap_or("-"),
1159 body.len(),
1160 started.elapsed(),
1161 );
1162 completion_transport_error(model, &action, err)
1163 })?;
1164 if body.len() + chunk.len() > MAX_COMPLETION_RESPONSE_BYTES {
1165 return Err(completion_response_too_large(
1166 model,
1167 request_id.as_deref(),
1168 body.len() + chunk.len(),
1169 ));
1170 }
1171 body.extend_from_slice(&chunk);
1172 let start = scanned.saturating_sub(SSE_DONE_MARKER.len());
1175 if body_contains_sse_done(&body, start) {
1176 return Ok((parse_streaming_json_events(&body, model)?, true));
1177 }
1178 scanned = body.len();
1179 }
1180
1181 Ok((parse_streaming_json_events(&body, model)?, false))
1182}
1183
1184#[cfg(test)]
1185async fn read_sse_json_events<T: DeserializeOwned>(
1186 response: reqwest::Response,
1187 model: &str,
1188) -> Result<Vec<T>, BoxError> {
1189 read_completion_stream(response, model)
1190 .await
1191 .map(|(events, _)| events)
1192}
1193
1194fn body_contains_sse_done(body: &[u8], from: usize) -> bool {
1200 if from == 0 && body.starts_with(SSE_DONE_MARKER) {
1201 return true;
1202 }
1203 body[from..]
1204 .windows(SSE_DONE_MARKER.len() + 1)
1205 .any(|window| window[0] == b'\n' && &window[1..] == SSE_DONE_MARKER)
1206}
1207
1208fn parse_streaming_json_events<T>(body: &[u8], model: &str) -> Result<Vec<T>, BoxError>
1209where
1210 T: DeserializeOwned,
1211{
1212 let body = std::str::from_utf8(body).map_err(|err| {
1213 format!(
1214 "Invalid UTF-8 in streaming completion response, model: {}, error: {}",
1215 model, err
1216 )
1217 })?;
1218 let body = body.strip_prefix('\u{feff}').unwrap_or(body);
1219
1220 if !looks_like_sse(body) {
1221 return parse_json_event_payload(body, model);
1222 }
1223
1224 let mut data = String::new();
1225 let mut events = Vec::new();
1226
1227 for line in body.lines() {
1228 let line = line.strip_suffix('\r').unwrap_or(line);
1229 handle_sse_text_line(line, &mut data, &mut events, model)?;
1230 }
1231 flush_sse_data(&mut data, &mut events, model)?;
1232
1233 Ok(events)
1234}
1235
1236fn looks_like_sse(body: &str) -> bool {
1237 body.lines().any(|line| {
1238 let line = line.strip_prefix('\u{feff}').unwrap_or(line);
1239 line.starts_with("data:")
1240 || line.starts_with("event:")
1241 || line.starts_with("id:")
1242 || line.starts_with("retry:")
1243 || line.starts_with(':')
1244 })
1245}
1246
1247fn handle_sse_text_line<T>(
1248 line: &str,
1249 data: &mut String,
1250 events: &mut Vec<T>,
1251 model: &str,
1252) -> Result<(), BoxError>
1253where
1254 T: DeserializeOwned,
1255{
1256 if line.is_empty() {
1257 return flush_sse_data(data, events, model);
1258 }
1259 if line.starts_with(':') {
1260 return Ok(());
1261 }
1262
1263 let Some(value) = line.strip_prefix("data:") else {
1264 return Ok(());
1265 };
1266 let value = value.strip_prefix(' ').unwrap_or(value);
1267 if !data.is_empty() {
1268 data.push('\n');
1269 }
1270 data.push_str(value);
1271 Ok(())
1272}
1273
1274fn flush_sse_data<T>(data: &mut String, events: &mut Vec<T>, model: &str) -> Result<(), BoxError>
1275where
1276 T: DeserializeOwned,
1277{
1278 let value = data.trim_end();
1279 if value.is_empty() || value == "[DONE]" {
1280 data.clear();
1281 return Ok(());
1282 }
1283
1284 let event = serde_json::from_str::<T>(value).map_err(|err| {
1285 format!(
1286 "Invalid streaming completion event, model: {}, error: {}, body: {}",
1287 model, err, value
1288 )
1289 })?;
1290 events.push(event);
1291 data.clear();
1292 Ok(())
1293}
1294
1295fn parse_json_event_payload<T>(body: &str, model: &str) -> Result<Vec<T>, BoxError>
1296where
1297 T: DeserializeOwned,
1298{
1299 let value = body.trim().strip_prefix('\u{feff}').unwrap_or(body.trim());
1300 if value.is_empty() || value == "[DONE]" {
1301 return Ok(Vec::new());
1302 }
1303
1304 if value.starts_with('[')
1305 && let Ok(events) = serde_json::from_str::<Vec<T>>(value)
1306 {
1307 return Ok(events);
1308 }
1309
1310 match serde_json::from_str::<T>(value) {
1311 Ok(event) => Ok(vec![event]),
1312 Err(single_err) => match serde_json::from_str::<Vec<T>>(value) {
1313 Ok(events) => Ok(events),
1314 Err(array_err) => {
1315 let mut events = Vec::new();
1316 let mut saw_line = false;
1317 for line in value.lines() {
1318 let line = line.trim();
1319 if line.is_empty() || line == "[DONE]" {
1320 continue;
1321 }
1322 saw_line = true;
1323 let event = serde_json::from_str::<T>(line).map_err(|line_err| {
1324 format!(
1325 "Invalid streaming completion event, model: {}, error: {}, body: {}",
1326 model, line_err, line
1327 )
1328 })?;
1329 events.push(event);
1330 }
1331
1332 if saw_line {
1333 return Ok(events);
1334 }
1335
1336 Err(format!(
1337 "Invalid streaming completion event, model: {}, error: {}; array error: {}, body: {}",
1338 model, single_err, array_err, value
1339 )
1340 .into())
1341 }
1342 },
1343 }
1344}
1345
1346pub(crate) fn streaming_completion_request(
1347 request: reqwest::RequestBuilder,
1348) -> reqwest::RequestBuilder {
1349 request
1350 .header(reqwest::header::ACCEPT, "text/event-stream")
1351 .header(reqwest::header::ACCEPT_ENCODING, "identity")
1352 .timeout(COMPLETION_REQUEST_TIMEOUT)
1357}
1358
1359#[cfg(test)]
1360mod tests {
1361 use super::*;
1362 use anda_core::FunctionDefinition;
1363 use http::{HeaderMap, HeaderValue, StatusCode};
1364 use tokio::io::{AsyncReadExt, AsyncWriteExt};
1365
1366 #[tokio::test]
1367 async fn prematurely_closed_provider_streams_are_retried() {
1368 use serde_json::json;
1369 let chat = json!({"id":"r","model":"test","choices":[{"index":0,"delta":{"role":"assistant","content":"ok"}}]});
1370 let gemini = json!({"candidates":[{"content":{"role":"model","parts":[{"text":"ok"}]}}]});
1371 let anthropic = json!({"type":"message_start","message":{"id":"r","type":"message","role":"assistant","model":"test","content":[{"type":"text","text":"ok"}],"usage":{}}});
1372 let cases = [
1373 (
1374 "openai",
1375 vec![chat],
1376 vec![json!({"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]})],
1377 ),
1378 (
1379 "gemini",
1380 vec![gemini],
1381 vec![json!({"candidates":[{"finishReason":"STOP"}]})],
1382 ),
1383 (
1384 "anthropic",
1385 vec![anthropic],
1386 vec![
1387 json!({"type":"message_delta","delta":{"stop_reason":"end_turn"}}),
1388 json!({"type":"message_stop"}),
1389 ],
1390 ),
1391 ];
1392 for (family, initial, terminal) in cases {
1393 let encode = |events: Vec<Json>| {
1394 events
1395 .into_iter()
1396 .map(|event| format!("data: {event}\n\n"))
1397 .collect::<String>()
1398 .into_bytes()
1399 };
1400 let incomplete = test_support::MockResponse {
1401 status: StatusCode::OK,
1402 headers: test_support::sse_headers(),
1403 body: encode(initial.clone()),
1404 };
1405 let complete = test_support::MockResponse {
1406 body: encode(initial.into_iter().chain(terminal).collect()),
1407 ..incomplete.clone()
1408 };
1409 let (endpoint, state) =
1410 test_support::spawn_retry_mock_server(vec![incomplete, complete]).await;
1411 let config = ModelConfig {
1412 family: family.into(),
1413 model: "test".into(),
1414 api_base: endpoint,
1415 api_key: "fake".into(),
1416 stream: true,
1417 ..Default::default()
1418 };
1419 let model = config.model(test_support::no_proxy_client()).unwrap();
1420 let output = model
1421 .completion(CompletionRequest {
1422 prompt: "hello".into(),
1423 ..Default::default()
1424 })
1425 .await
1426 .unwrap();
1427 assert_eq!(output.content, "ok", "{family}");
1428 assert!(output.failed_reason.is_none(), "{family}");
1429 assert_eq!(state.lock().unwrap().1, 2, "{family}");
1430 }
1431 }
1432
1433 #[derive(Clone)]
1434 struct TestCompleter {
1435 name: &'static str,
1436 }
1437
1438 impl CompletionFeaturesDyn for TestCompleter {
1439 fn completion(&self, _req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>> {
1440 Box::pin(futures::future::ready(Ok(AgentOutput::default())))
1441 }
1442
1443 fn model_name(&self) -> String {
1444 self.name.to_string()
1445 }
1446 }
1447
1448 fn test_model(name: &'static str) -> Model {
1449 Model::new(Arc::new(TestCompleter { name }))
1450 }
1451
1452 fn http_client() -> reqwest::Client {
1453 test_support::no_proxy_client()
1454 }
1455
1456 async fn spawn_truncated_sse_after_done_server() -> String {
1457 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1458 let addr = listener.local_addr().unwrap();
1459 tokio::spawn(async move {
1460 let (mut socket, _) = listener.accept().await.unwrap();
1461 let mut request = [0; 1024];
1462 let _ = socket.read(&mut request).await;
1463 socket
1464 .write_all(
1465 b"HTTP/1.1 200 OK\r\n\
1466 Content-Type: text/event-stream\r\n\
1467 Content-Length: 4096\r\n\
1468 Connection: close\r\n\
1469 \r\n\
1470 data: {\"a\":1}\n\n\
1471 data: [DONE]\n\n",
1472 )
1473 .await
1474 .unwrap();
1475 let _ = socket.shutdown().await;
1476 });
1477 format!("http://{addr}")
1478 }
1479
1480 async fn spawn_sse_with_done_marker_in_content_server() -> String {
1484 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1485 let addr = listener.local_addr().unwrap();
1486 tokio::spawn(async move {
1487 let (mut socket, _) = listener.accept().await.unwrap();
1488 let mut request = [0; 1024];
1489 let _ = socket.read(&mut request).await;
1490 socket
1491 .write_all(
1492 b"HTTP/1.1 200 OK\r\n\
1493 Content-Type: text/event-stream\r\n\
1494 Connection: close\r\n\
1495 \r\n\
1496 data: {\"text\":\"sse ends with data: [DONE]\"}\n\n",
1497 )
1498 .await
1499 .unwrap();
1500 socket.flush().await.unwrap();
1501 tokio::time::sleep(Duration::from_millis(50)).await;
1502 socket
1503 .write_all(b"data: {\"b\":2}\n\ndata: [DONE]\n\n")
1504 .await
1505 .unwrap();
1506 let _ = socket.shutdown().await;
1507 });
1508 format!("http://{addr}")
1509 }
1510
1511 async fn spawn_stalling_sse_body_server() -> String {
1512 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1513 let addr = listener.local_addr().unwrap();
1514 tokio::spawn(async move {
1515 let (mut socket, _) = listener.accept().await.unwrap();
1516 let mut request = [0; 1024];
1517 let _ = socket.read(&mut request).await;
1518 socket
1519 .write_all(
1520 b"HTTP/1.1 200 OK\r\n\
1521 Content-Type: text/event-stream\r\n\
1522 Connection: close\r\n\
1523 \r\n\
1524 data: {\"a\":1}\n\n",
1525 )
1526 .await
1527 .unwrap();
1528 socket.flush().await.unwrap();
1529 tokio::time::sleep(Duration::from_millis(500)).await;
1530 let _ = socket.shutdown().await;
1531 });
1532 format!("http://{addr}")
1533 }
1534
1535 fn retry_count(state: &test_support::RetryState) -> usize {
1536 state.lock().unwrap().1
1537 }
1538
1539 fn model_config(family: &str, model: &str) -> ModelConfig {
1540 ModelConfig {
1541 family: family.to_string(),
1542 model: model.to_string(),
1543 api_base: "https://example.com".to_string(),
1544 api_key: "test-key".to_string(),
1545 ..Default::default()
1546 }
1547 }
1548
1549 #[test]
1550 fn model_effort_serializes_config_values() {
1551 let config: ModelConfig = serde_json::from_value(serde_json::json!({
1552 "family": "openai",
1553 "model": "gpt-5",
1554 "api_base": "http://localhost",
1555 "api_key": "test-key",
1556 "effort": "max"
1557 }))
1558 .unwrap();
1559
1560 assert_eq!(config.effort, Some(ModelEffort::Max));
1561 assert_eq!(
1562 serde_json::to_value(ModelEffort::Minimal).unwrap(),
1563 "minimal"
1564 );
1565 }
1566
1567 #[test]
1568 fn models_default_is_empty() {
1569 let models = Models::default();
1570
1571 assert!(models.get_model().is_none());
1572 assert!(models.get("missing").is_none());
1573 assert!(models.resolve("missing").is_none());
1574 }
1575
1576 #[test]
1577 fn set_model_sets_primary_without_registering_a_label() {
1578 let models = Models::default();
1579 models.set_model(test_model("primary"));
1580
1581 assert_eq!(
1582 models
1583 .get_model()
1584 .expect("primary model should exist")
1585 .model_name(),
1586 "primary"
1587 );
1588 assert!(models.get("primary").is_some());
1589 }
1590
1591 #[test]
1592 fn set_promotes_first_inserted_model_to_primary() {
1593 let models = Models::default();
1594 models.set("x".to_string(), test_model("X"));
1595
1596 assert_eq!(
1597 models.get("x").expect("label x should exist").model_name(),
1598 "X"
1599 );
1600 assert_eq!(
1601 models
1602 .get_model()
1603 .expect("primary model should be initialized")
1604 .model_name(),
1605 "X"
1606 );
1607 }
1608
1609 #[test]
1610 fn fallback_label_has_no_default_routing_semantics() {
1611 let models = Models::default();
1612 models.set_model(test_model("primary"));
1613 models.set("fallback".to_string(), test_model("fallback"));
1614
1615 assert_eq!(
1616 models
1617 .get("fallback")
1618 .expect("fallback is still a normal label")
1619 .model_name(),
1620 "fallback"
1621 );
1622 assert_eq!(
1623 models
1624 .get_model()
1625 .expect("primary model should stay the default")
1626 .model_name(),
1627 "primary"
1628 );
1629 assert_eq!(
1630 models
1631 .resolve("unknown")
1632 .expect("missing label should use default routing")
1633 .model_name(),
1634 "primary"
1635 );
1636 }
1637
1638 #[test]
1639 fn resolve_prefers_exact_label_then_default() {
1640 let models = Models::default();
1641 models.set_model(test_model("primary"));
1642 models.set("flash".to_string(), test_model("flash"));
1643
1644 assert_eq!(
1645 models
1646 .resolve("flash")
1647 .expect("exact label should win")
1648 .model_name(),
1649 "flash"
1650 );
1651 assert_eq!(
1652 models
1653 .resolve("missing")
1654 .expect("missing label should use default routing")
1655 .model_name(),
1656 "primary"
1657 );
1658 assert_eq!(
1659 models
1660 .resolve("")
1661 .expect("empty label should use default routing")
1662 .model_name(),
1663 "primary"
1664 );
1665 }
1666
1667 #[test]
1668 fn model_config_validates_required_fields_and_builds_supported_families() {
1669 let client = http_client();
1670
1671 let mut config = model_config("openai", "gpt-5");
1672 config.disabled = true;
1673 let Err(err) = config.model(client.clone()) else {
1674 panic!("disabled model should fail");
1675 };
1676 assert!(err.to_string().contains("disabled"));
1677
1678 for (field, config) in [
1679 (
1680 "model name",
1681 ModelConfig {
1682 model: String::new(),
1683 ..model_config("openai", "gpt-5")
1684 },
1685 ),
1686 (
1687 "model family",
1688 ModelConfig {
1689 family: String::new(),
1690 ..model_config("openai", "gpt-5")
1691 },
1692 ),
1693 (
1694 "api_base",
1695 ModelConfig {
1696 api_base: String::new(),
1697 ..model_config("openai", "gpt-5")
1698 },
1699 ),
1700 (
1701 "api_key",
1702 ModelConfig {
1703 api_key: String::new(),
1704 ..model_config("openai", "gpt-5")
1705 },
1706 ),
1707 ] {
1708 let Err(err) = config.model(client.clone()) else {
1709 panic!("{field} should fail");
1710 };
1711 let err = err.to_string();
1712 assert!(err.contains(field), "{field}: {err}");
1713 }
1714
1715 let Err(err) = model_config("unknown", "m").model(client.clone()) else {
1716 panic!("unsupported family should fail");
1717 };
1718 assert!(err.to_string().contains("unsupported model family"));
1719
1720 let mut gemini = model_config("gemini", "gemini-2.5-pro");
1721 gemini.context_window = 123;
1722 gemini.max_output = 45;
1723 let model = gemini.model(client.clone()).unwrap();
1724 assert_eq!(model.model_name(), "gemini-2.5-pro");
1725 assert_eq!(model.labels, vec!["gemini-2.5-pro"]);
1726 assert_eq!(model.context_window, 123);
1727 assert_eq!(model.max_output, 45);
1728
1729 let mut anthropic = model_config("anthropic", "claude-sonnet-4-5");
1730 anthropic.labels = vec!["pro".to_string(), "primary".to_string()];
1731 anthropic.bearer_auth = true;
1732 anthropic.stream = true;
1733 anthropic.effort = Some(ModelEffort::High);
1734 let model = anthropic.model(client.clone()).unwrap();
1735 assert_eq!(model.model_name(), "claude-sonnet-4-5");
1736 assert_eq!(model.labels, vec!["pro", "primary"]);
1737
1738 let model = model_config("openai", "gpt-5")
1739 .model(client.clone())
1740 .unwrap();
1741 assert_eq!(model.model_name(), "gpt-5");
1742 let model = model_config("openai", "deepseek-chat")
1743 .model(client)
1744 .unwrap();
1745 assert_eq!(model.model_name(), "deepseek-chat");
1746 }
1747
1748 #[test]
1749 fn models_registry_clones_names_replaces_labels_and_loads_configs() {
1750 let models = Models::default();
1751 models.set_model(test_model("flash-v1").with_labels(vec!["FAST".into()]));
1752 assert!(models.contains("fast"));
1753 assert_eq!(
1754 models.model_names(),
1755 BTreeSet::from(["flash-v1".to_string()])
1756 );
1757
1758 models.set("flash".to_string(), test_model("flash-v2"));
1759 assert!(models.contains("flash"));
1760 assert_eq!(models.get("FLASH").unwrap().model_name(), "flash-v2");
1761 assert_eq!(
1762 models.model_names(),
1763 BTreeSet::from(["flash-v1".to_string(), "flash-v2".to_string()])
1764 );
1765
1766 models.set("primary".to_string(), test_model("primary-v2"));
1767 assert_eq!(models.get_model().unwrap().model_name(), "primary-v2");
1768
1769 let cloned = Models::from_clone(&models);
1770 assert_eq!(cloned.get("primary").unwrap().model_name(), "primary-v2");
1771 assert_eq!(
1772 cloned.resolve("missing").unwrap().model_name(),
1773 "primary-v2"
1774 );
1775
1776 let replacement = Models::default();
1777 replacement.set_model(test_model("replacement-primary").with_labels(vec!["next".into()]));
1778 let replaced = Models::default();
1779 replaced.set("old".to_string(), test_model("old"));
1780 replaced.replace(&replacement);
1781 assert!(!replaced.contains("old"));
1782 assert!(replaced.contains("next"));
1783 assert_eq!(
1784 replaced.get_model().unwrap().model_name(),
1785 "replacement-primary"
1786 );
1787
1788 replacement.set("later".to_string(), test_model("later"));
1789 assert!(!replaced.contains("later"));
1790
1791 let configs = vec![
1792 ModelConfig {
1793 labels: vec!["primary".to_string()],
1794 ..model_config("openai", "gpt-5")
1795 },
1796 ModelConfig {
1797 disabled: true,
1798 ..model_config("openai", "disabled")
1799 },
1800 ];
1801 let loaded = Models::from_configs(&configs, http_client());
1802 assert!(loaded.contains("primary"));
1803 assert!(!loaded.contains("disabled"));
1804 assert_eq!(loaded.get_model().unwrap().model_name(), "gpt-5");
1805 }
1806
1807 #[tokio::test]
1808 async fn model_max_output_caps_explicit_request_output_budget() {
1809 let completer = testing::ScriptedCompleter::new("capped").into_arc();
1810 let mut model = Model::with_completer(completer.clone());
1811
1812 model
1814 .completion(CompletionRequest::default())
1815 .await
1816 .unwrap();
1817
1818 model.max_output = 32_000;
1819 model
1820 .completion(CompletionRequest::default())
1821 .await
1822 .unwrap();
1823 for requested in [8_000, 64_000] {
1824 model
1825 .completion(CompletionRequest {
1826 max_output_tokens: Some(requested),
1827 ..Default::default()
1828 })
1829 .await
1830 .unwrap();
1831 }
1832
1833 let sent = completer
1834 .requests()
1835 .into_iter()
1836 .map(|req| req.max_output_tokens)
1837 .collect::<Vec<_>>();
1838 assert_eq!(sent, vec![None, None, Some(8_000), Some(32_000)]);
1840 }
1841
1842 #[tokio::test]
1843 async fn model_completion_placeholders_and_mock_tool_calls_are_stable() {
1844 let not_implemented = Model::not_implemented();
1845 assert_eq!(not_implemented.model_name(), "not_implemented");
1846 let err = not_implemented
1847 .completion(CompletionRequest::default())
1848 .await
1849 .unwrap_err();
1850 assert!(err.to_string().contains("not implemented"));
1851
1852 let mock = Model::mock_implemented().with_labels(vec!["mock".into()]);
1853 assert_eq!(mock.model_name(), "mock_implemented");
1854 let output = mock
1855 .completion(CompletionRequest {
1856 prompt: "{\"q\":\"anda\"}".to_string(),
1857 tools: vec![FunctionDefinition {
1858 name: "lookup".to_string(),
1859 ..Default::default()
1860 }],
1861 ..Default::default()
1862 })
1863 .await
1864 .unwrap();
1865 assert_eq!(output.content, "{\"q\":\"anda\"}");
1866 assert_eq!(output.tool_calls.len(), 1);
1867 assert_eq!(output.tool_calls[0].name, "lookup");
1868 assert_eq!(output.tool_calls[0].args["q"], "anda");
1869
1870 let output = mock
1871 .completion(CompletionRequest {
1872 prompt: String::new(),
1873 tools: vec![FunctionDefinition {
1874 name: "lookup".to_string(),
1875 ..Default::default()
1876 }],
1877 ..Default::default()
1878 })
1879 .await
1880 .unwrap();
1881 assert!(output.tool_calls.is_empty());
1882 }
1883
1884 #[test]
1885 fn streaming_json_event_parser_accepts_bom_sse_ndjson_and_arrays() {
1886 let events = parse_streaming_json_events::<serde_json::Value>(
1887 b"\xef\xbb\xbfdata: {\"a\":1}\n\ndata: [DONE]\n\n",
1888 "test-model",
1889 )
1890 .unwrap();
1891 assert_eq!(events, vec![serde_json::json!({"a": 1})]);
1892
1893 let events = parse_streaming_json_events::<serde_json::Value>(
1894 b"{\"a\":1}\n{\"b\":2}\n[DONE]\n",
1895 "test-model",
1896 )
1897 .unwrap();
1898 assert_eq!(
1899 events,
1900 vec![serde_json::json!({"a": 1}), serde_json::json!({"b": 2})]
1901 );
1902
1903 let events =
1904 parse_streaming_json_events::<serde_json::Value>(br#"[{"a":1},{"b":2}]"#, "test-model")
1905 .unwrap();
1906 assert_eq!(
1907 events,
1908 vec![serde_json::json!({"a": 1}), serde_json::json!({"b": 2})]
1909 );
1910 }
1911
1912 #[tokio::test]
1913 async fn streaming_reader_ignores_mislabelled_content_encoding() {
1914 let mut headers = HeaderMap::new();
1915 headers.insert(
1916 http::header::CONTENT_TYPE,
1917 HeaderValue::from_static("text/event-stream"),
1918 );
1919 let (endpoint, _) =
1920 test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
1921 status: StatusCode::OK,
1922 headers,
1923 body: b"data: {\"a\":1}\n\ndata: [DONE]\n\n".to_vec(),
1924 }])
1925 .await;
1926 let client = request_client_builder()
1927 .https_only(false)
1928 .no_proxy()
1929 .build()
1930 .unwrap();
1931 let response = client.get(endpoint).send().await.unwrap();
1932
1933 let events = read_sse_json_events::<serde_json::Value>(response, "test-model")
1934 .await
1935 .unwrap();
1936
1937 assert_eq!(events, vec![serde_json::json!({"a": 1})]);
1938 }
1939
1940 #[tokio::test]
1941 async fn streaming_reader_returns_after_done_before_late_body_error() {
1942 let endpoint = spawn_truncated_sse_after_done_server().await;
1943 let client = request_client_builder()
1944 .https_only(false)
1945 .no_proxy()
1946 .build()
1947 .unwrap();
1948 let response = client.get(endpoint).send().await.unwrap();
1949
1950 let events = read_sse_json_events::<serde_json::Value>(response, "test-model")
1951 .await
1952 .unwrap();
1953
1954 assert_eq!(events, vec![serde_json::json!({"a": 1})]);
1955 }
1956
1957 #[test]
1958 fn sse_done_detection_is_line_anchored() {
1959 assert!(body_contains_sse_done(b"data: [DONE]\n\n", 0));
1960 assert!(body_contains_sse_done(
1961 b"data: {\"a\":1}\n\ndata: [DONE]\n\n",
1962 0
1963 ));
1964 assert!(!body_contains_sse_done(
1967 b"data: {\"text\":\"sse ends with data: [DONE]\"}\n\n",
1968 0
1969 ));
1970 }
1971
1972 #[tokio::test]
1973 async fn streaming_reader_is_not_truncated_by_done_marker_in_content() {
1974 let endpoint = spawn_sse_with_done_marker_in_content_server().await;
1975 let client = request_client_builder()
1976 .https_only(false)
1977 .no_proxy()
1978 .build()
1979 .unwrap();
1980 let response = client.get(endpoint).send().await.unwrap();
1981
1982 let events = read_sse_json_events::<serde_json::Value>(response, "test-model")
1983 .await
1984 .unwrap();
1985
1986 assert_eq!(
1987 events,
1988 vec![
1989 serde_json::json!({"text": "sse ends with data: [DONE]"}),
1990 serde_json::json!({"b": 2})
1991 ]
1992 );
1993 }
1994
1995 #[test]
1996 fn completion_transport_timeouts_are_streaming_safe() {
1997 assert_eq!(COMPLETION_HTTP2_KEEP_ALIVE_INTERVAL, None);
2001 assert!(COMPLETION_READ_TIMEOUT > Duration::from_secs(118));
2002 assert!(COMPLETION_READ_TIMEOUT < COMPLETION_REQUEST_TIMEOUT);
2003 assert_eq!(COMPLETION_REQUEST_TIMEOUT, Duration::from_secs(600));
2004 }
2005
2006 #[test]
2007 fn streaming_completion_request_overrides_short_client_total_timeout() {
2008 let client = reqwest::Client::builder()
2009 .no_proxy()
2010 .timeout(Duration::from_millis(100))
2011 .build()
2012 .unwrap();
2013 let request = streaming_completion_request(client.get("https://example.com"))
2014 .build()
2015 .unwrap();
2016
2017 assert_eq!(request.timeout(), Some(&COMPLETION_REQUEST_TIMEOUT));
2018 }
2019
2020 #[tokio::test]
2021 async fn streaming_reader_body_idle_timeout_is_retryable() {
2022 let endpoint = spawn_stalling_sse_body_server().await;
2023 let client = request_client_builder()
2024 .https_only(false)
2025 .no_proxy()
2026 .read_timeout(Duration::from_millis(100))
2027 .timeout(Duration::from_secs(5))
2028 .build()
2029 .unwrap();
2030 let response = client.get(endpoint).send().await.unwrap();
2031
2032 let err = tokio::time::timeout(
2033 Duration::from_secs(2),
2034 read_sse_json_events::<serde_json::Value>(response, "test-model"),
2035 )
2036 .await
2037 .expect("body read timeout should fire")
2038 .unwrap_err();
2039
2040 let message = err.to_string();
2041 assert!(
2042 message.contains("Failed to read streaming completion response"),
2043 "{message}"
2044 );
2045 assert!(message.contains("received:"), "{message}");
2046 assert!(message.contains("operation timed out"), "{message}");
2047 assert!(is_retryable_box_error(&err));
2048 }
2049
2050 #[test]
2051 fn retry_after_parses_seconds_and_http_date() {
2052 let mut headers = HeaderMap::new();
2053 headers.insert(http::header::RETRY_AFTER, HeaderValue::from_static("42"));
2054 assert_eq!(
2055 retry_after_duration(&headers),
2056 Some(Duration::from_secs(42))
2057 );
2058
2059 let when = chrono::Utc::now() + chrono::Duration::seconds(90);
2060 headers.insert(
2061 http::header::RETRY_AFTER,
2062 HeaderValue::from_str(&when.to_rfc2822()).unwrap(),
2063 );
2064 let parsed = retry_after_duration(&headers).expect("http-date should parse");
2065 assert!(parsed <= Duration::from_secs(90));
2066 assert!(parsed >= Duration::from_secs(80));
2067
2068 let when = chrono::Utc::now() - chrono::Duration::seconds(90);
2070 headers.insert(
2071 http::header::RETRY_AFTER,
2072 HeaderValue::from_str(&when.to_rfc2822()).unwrap(),
2073 );
2074 assert_eq!(retry_after_duration(&headers), None);
2075
2076 headers.insert(
2077 http::header::RETRY_AFTER,
2078 HeaderValue::from_static("not-a-date"),
2079 );
2080 assert_eq!(retry_after_duration(&headers), None);
2081 }
2082
2083 #[tokio::test]
2084 async fn custom_client_streaming_decode_errors_are_retryable() {
2085 let mut headers = HeaderMap::new();
2086 headers.insert(
2087 http::header::CONTENT_TYPE,
2088 HeaderValue::from_static("text/event-stream"),
2089 );
2090 headers.insert(
2091 http::header::CONTENT_ENCODING,
2092 HeaderValue::from_static("gzip"),
2093 );
2094 let (endpoint, _) =
2095 test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
2096 status: StatusCode::OK,
2097 headers,
2098 body: b"data: {\"a\":1}\n\ndata: [DONE]\n\n".to_vec(),
2099 }])
2100 .await;
2101 let client = reqwest::Client::builder().no_proxy().build().unwrap();
2102 let response = streaming_completion_request(client.get(endpoint))
2103 .send()
2104 .await
2105 .unwrap();
2106
2107 let err = read_sse_json_events::<serde_json::Value>(response, "test-model")
2108 .await
2109 .unwrap_err();
2110
2111 let message = err.to_string();
2112 assert!(message.contains("error decoding response body"));
2113 assert!(message.contains("received: 0 bytes"), "{message}");
2116 assert!(message.contains("request_id: -"), "{message}");
2117 assert!(
2118 message.contains("error decoding response body: "),
2119 "{message}"
2120 );
2121 assert!(is_retryable_box_error(&err));
2122 }
2123
2124 #[test]
2125 fn error_chain_formatting_appends_unique_sources() {
2126 let root = std::io::Error::new(std::io::ErrorKind::TimedOut, "operation timed out");
2127 let outer =
2128 ModelError::new("error decoding response body".to_string()).with_source(Box::new(root));
2129 assert_eq!(
2130 format_error_chain(&outer),
2131 "error decoding response body: operation timed out"
2132 );
2133
2134 let root = std::io::Error::new(std::io::ErrorKind::TimedOut, "operation timed out");
2136 let outer = ModelError::new("request failed: operation timed out".to_string())
2137 .with_source(Box::new(root));
2138 assert_eq!(
2139 format_error_chain(&outer),
2140 "request failed: operation timed out"
2141 );
2142 }
2143
2144 #[tokio::test]
2145 async fn completion_error_bodies_are_truncated_for_diagnostics() {
2146 assert_eq!(error_body_excerpt(b"short body"), "short body");
2147 let excerpt = error_body_excerpt(&vec![b'x'; MAX_ERROR_BODY_BYTES * 4]);
2148 assert!(excerpt.ends_with("… [truncated]"));
2149 assert!(excerpt.len() < MAX_ERROR_BODY_BYTES + 32);
2150
2151 let (endpoint, _) =
2152 test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
2153 status: StatusCode::BAD_REQUEST,
2154 headers: HeaderMap::new(),
2155 body: vec![b'e'; 1024 * 1024],
2156 }])
2157 .await;
2158 let client = http_client();
2159 let err = execute_completion_request_with_retry(
2160 "error-body-test",
2161 || client.post(&endpoint),
2162 |response| async { read_completion_response_bytes(response, "error-body-test").await },
2163 )
2164 .await
2165 .unwrap_err();
2166 let message = err.to_string();
2167 assert!(message.contains("status: 400"));
2168 assert!(message.ends_with("… [truncated]"));
2169 assert!(message.len() < MAX_ERROR_BODY_BYTES + 256);
2170 }
2171
2172 #[test]
2173 fn upstream_request_id_checks_known_headers() {
2174 let mut headers = HeaderMap::new();
2175 assert_eq!(upstream_request_id(&headers), None);
2176
2177 headers.insert("cf-ray", HeaderValue::from_static("ray-123"));
2178 assert_eq!(upstream_request_id(&headers), Some("ray-123".to_string()));
2179
2180 headers.insert("x-request-id", HeaderValue::from_static("req-456"));
2181 assert_eq!(upstream_request_id(&headers), Some("req-456".to_string()));
2182 }
2183
2184 #[tokio::test]
2185 async fn completion_request_retries_transient_errors_and_exposes_retry_signal() {
2186 let mut headers = HeaderMap::new();
2187 headers.insert(http::header::RETRY_AFTER, HeaderValue::from_static("0"));
2188 let (endpoint, state) = test_support::spawn_retry_mock_server(vec![
2189 test_support::MockResponse {
2190 status: StatusCode::TOO_MANY_REQUESTS,
2191 headers,
2192 body: b"rate limited".to_vec(),
2193 },
2194 test_support::MockResponse {
2195 status: StatusCode::OK,
2196 headers: HeaderMap::new(),
2197 body: b"ok".to_vec(),
2198 },
2199 ])
2200 .await;
2201 let client = http_client();
2202
2203 let body = execute_completion_request_with_retry(
2204 "retry-test",
2205 || client.post(&endpoint),
2206 |response| async { read_completion_response_bytes(response, "retry-test").await },
2207 )
2208 .await
2209 .unwrap();
2210
2211 assert_eq!(&body[..], b"ok");
2212 assert_eq!(retry_count(&state), 2);
2213
2214 let mut retry_now = HeaderMap::new();
2215 retry_now.insert(http::header::RETRY_AFTER, HeaderValue::from_static("0"));
2216 let mut final_headers = HeaderMap::new();
2217 final_headers.insert(http::header::RETRY_AFTER, HeaderValue::from_static("45"));
2218 let (endpoint, state) = test_support::spawn_retry_mock_server(vec![
2219 test_support::MockResponse {
2220 status: StatusCode::TOO_MANY_REQUESTS,
2221 headers: retry_now.clone(),
2222 body: b"first limit".to_vec(),
2223 },
2224 test_support::MockResponse {
2225 status: StatusCode::TOO_MANY_REQUESTS,
2226 headers: retry_now.clone(),
2227 body: b"still limited".to_vec(),
2228 },
2229 test_support::MockResponse {
2230 status: StatusCode::TOO_MANY_REQUESTS,
2231 headers: retry_now,
2232 body: b"limited again".to_vec(),
2233 },
2234 test_support::MockResponse {
2235 status: StatusCode::TOO_MANY_REQUESTS,
2236 headers: final_headers,
2237 body: b"finally limited".to_vec(),
2238 },
2239 ])
2240 .await;
2241 let err = execute_completion_request_with_retry(
2242 "retry-test",
2243 || client.post(&endpoint),
2244 |response| async { read_completion_response_bytes(response, "retry-test").await },
2245 )
2246 .await
2247 .unwrap_err();
2248 let err_ref = err.as_ref() as &(dyn Error + 'static);
2249
2250 assert_eq!(retry_count(&state), MODEL_REQUEST_MAX_RETRIES + 1);
2251 assert!(is_retryable_box_error(&err));
2252 assert_eq!(
2253 model_error_status(err_ref),
2254 Some(StatusCode::TOO_MANY_REQUESTS)
2255 );
2256 assert_eq!(
2257 model_error_retry_after(err_ref),
2258 Some(Duration::from_secs(45))
2259 );
2260
2261 let (endpoint, state) =
2262 test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
2263 status: StatusCode::BAD_REQUEST,
2264 headers: HeaderMap::new(),
2265 body: b"bad request".to_vec(),
2266 }])
2267 .await;
2268 let err = execute_completion_request_with_retry(
2269 "retry-test",
2270 || client.post(&endpoint),
2271 |response| async { read_completion_response_bytes(response, "retry-test").await },
2272 )
2273 .await
2274 .unwrap_err();
2275
2276 assert_eq!(retry_count(&state), 1);
2277 assert!(!is_retryable_box_error(&err));
2278 }
2279}