starweaver_model/wrappers/
fallback.rs1use 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
18pub 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 #[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 #[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 #[must_use]
58 pub fn len(&self) -> usize {
59 self.models.len()
60 }
61
62 #[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}