Skip to main content

starweaver_model/wrappers/
fallback.rs

1//! Fallback model wrapper.
2
3use async_trait::async_trait;
4use serde_json::{Value, json};
5
6use super::DynModelAdapter;
7use crate::{
8    adapter::{
9        ModelAdapter, ModelError, ModelRequestContext, ModelRequestParameters,
10        ModelResponseEventStream,
11    },
12    message::{ModelMessage, ModelResponse},
13    profile::ModelProfile,
14    settings::ModelSettings,
15    stream::ModelResponseStreamEvent,
16};
17
18/// Model wrapper that tries adapters in order until one succeeds.
19pub struct FallbackModel {
20    models: Vec<DynModelAdapter>,
21    model_name: String,
22    provider_name: Option<String>,
23    profile: ModelProfile,
24    default_settings: Option<ModelSettings>,
25}
26
27impl FallbackModel {
28    /// Create a fallback wrapper from ordered candidate models.
29    ///
30    /// # Panics
31    ///
32    /// Panics when `models` is empty because a fallback wrapper requires at least one candidate.
33    #[must_use]
34    pub fn new(models: Vec<DynModelAdapter>) -> Self {
35        assert!(
36            !models.is_empty(),
37            "fallback model requires at least one candidate"
38        );
39        let primary = models[0].clone();
40        Self {
41            model_name: primary.model_name().to_string(),
42            provider_name: primary.provider_name().map(str::to_string),
43            profile: primary.profile().clone(),
44            default_settings: primary.default_settings().cloned(),
45            models,
46        }
47    }
48
49    /// Override the exposed model name.
50    #[must_use]
51    pub fn with_model_name(mut self, model_name: impl Into<String>) -> Self {
52        self.model_name = model_name.into();
53        self
54    }
55
56    /// Return candidate count.
57    #[must_use]
58    pub fn len(&self) -> usize {
59        self.models.len()
60    }
61
62    /// Return whether there are no candidates.
63    #[must_use]
64    pub fn is_empty(&self) -> bool {
65        self.models.is_empty()
66    }
67}
68
69#[async_trait]
70impl ModelAdapter for FallbackModel {
71    fn model_name(&self) -> &str {
72        &self.model_name
73    }
74
75    fn provider_name(&self) -> Option<&str> {
76        self.provider_name.as_deref()
77    }
78
79    fn profile(&self) -> &ModelProfile {
80        &self.profile
81    }
82
83    fn default_settings(&self) -> Option<&ModelSettings> {
84        self.default_settings.as_ref()
85    }
86
87    async fn request(
88        &self,
89        messages: Vec<ModelMessage>,
90        settings: Option<ModelSettings>,
91        params: ModelRequestParameters,
92        context: ModelRequestContext,
93    ) -> Result<ModelResponse, ModelError> {
94        let mut failures = Vec::new();
95        let mut attempts = 0u32;
96        for model in &self.models {
97            attempts += 1;
98            let mut attempt_context = context.clone();
99            annotate_attempt(&mut attempt_context, "request", attempts, model.as_ref());
100            match model
101                .request(
102                    messages.clone(),
103                    settings.clone(),
104                    params.clone(),
105                    attempt_context,
106                )
107                .await
108            {
109                Ok(mut response) => {
110                    annotate_response_success(&mut response, attempts, model.as_ref(), &failures);
111                    return Ok(response);
112                }
113                Err(error) => failures.push(fallback_failure(attempts, model.as_ref(), &error)),
114            }
115        }
116        Err(fallback_error(attempts, failures))
117    }
118
119    async fn request_stream(
120        &self,
121        messages: Vec<ModelMessage>,
122        settings: Option<ModelSettings>,
123        params: ModelRequestParameters,
124        context: ModelRequestContext,
125    ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
126        let mut failures = Vec::new();
127        let mut attempts = 0u32;
128        for model in &self.models {
129            attempts += 1;
130            let mut attempt_context = context.clone();
131            annotate_attempt(
132                &mut attempt_context,
133                "request_stream",
134                attempts,
135                model.as_ref(),
136            );
137            match model
138                .request_stream(
139                    messages.clone(),
140                    settings.clone(),
141                    params.clone(),
142                    attempt_context,
143                )
144                .await
145            {
146                Ok(mut events) => {
147                    annotate_stream_success(&mut events, attempts, model.as_ref(), &failures);
148                    return Ok(events);
149                }
150                Err(error) => failures.push(fallback_failure(attempts, model.as_ref(), &error)),
151            }
152        }
153        Err(fallback_error(attempts, failures))
154    }
155
156    async fn request_stream_incremental(
157        &self,
158        messages: Vec<ModelMessage>,
159        settings: Option<ModelSettings>,
160        params: ModelRequestParameters,
161        context: ModelRequestContext,
162    ) -> Result<ModelResponseEventStream, ModelError> {
163        let events = self
164            .request_stream(messages, settings, params, context)
165            .await?;
166        let (sender, receiver) = tokio::sync::mpsc::channel(events.len().max(1));
167        tokio::spawn(async move {
168            for event in events {
169                if sender.send(Ok(event)).await.is_err() {
170                    return;
171                }
172            }
173        });
174        Ok(ModelResponseEventStream::new(receiver))
175    }
176}
177
178fn annotate_attempt(
179    context: &mut ModelRequestContext,
180    call_kind: &str,
181    attempt: u32,
182    model: &dyn ModelAdapter,
183) {
184    context.llm_trace_metadata.insert(
185        "starweaver_model_wrapper".to_string(),
186        json!({
187            "kind": "fallback",
188            "call_kind": call_kind,
189            "attempt": attempt,
190            "model": model.model_name(),
191            "provider": model.provider_name(),
192        }),
193    );
194}
195
196fn annotate_response_success(
197    response: &mut ModelResponse,
198    selected_attempt: u32,
199    model: &dyn ModelAdapter,
200    failures: &[Value],
201) {
202    response.metadata.insert(
203        "starweaver_model_wrapper".to_string(),
204        json!({
205            "kind": "fallback",
206            "selected_attempt": selected_attempt,
207            "selected_model": model.model_name(),
208            "selected_provider": model.provider_name(),
209            "failures": failures,
210        }),
211    );
212}
213
214fn annotate_stream_success(
215    events: &mut [ModelResponseStreamEvent],
216    selected_attempt: u32,
217    model: &dyn ModelAdapter,
218    failures: &[Value],
219) {
220    for event in events.iter_mut() {
221        if let ModelResponseStreamEvent::FinalResult(response) = event {
222            annotate_response_success(response, selected_attempt, model, failures);
223        }
224    }
225}
226
227fn fallback_failure(attempt: u32, model: &dyn ModelAdapter, error: &ModelError) -> Value {
228    json!({
229        "attempt": attempt,
230        "model": model.model_name(),
231        "provider": model.provider_name(),
232        "error": error.to_string(),
233    })
234}
235
236fn fallback_error(attempts: u32, failures: Vec<Value>) -> ModelError {
237    ModelError::Transport(format!(
238        "fallback model exhausted after {attempts} attempts: {}",
239        Value::Array(failures)
240    ))
241}