1use std::path::{Path, PathBuf};
2use std::sync::Arc;
3
4use serde::{Deserialize, Deserializer, Serialize, Serializer};
5use thiserror::Error;
6
7use crate::config::{
8 build_vv_llm_from_local_settings, resolve_model_endpoint, ResolvedModelConfig,
9};
10use crate::llm::{LlmClient, LlmRequest, ScriptStep, ScriptedLlmClient, VvLlmClient};
11use crate::model_settings::ModelSettings;
12use crate::types::LLMResponse;
13
14#[derive(Debug, Clone, PartialEq, Eq)]
15pub enum ModelRef {
16 Named(String),
17 BackendModel { backend: String, model: String },
18 Resolved(ResolvedModelConfig),
19}
20
21#[derive(Serialize, Deserialize)]
22#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
23enum ModelRefWire {
24 Named { model: String },
25 BackendModel { backend: String, model: String },
26}
27
28impl Serialize for ModelRef {
29 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
30 where
31 S: Serializer,
32 {
33 let wire = match self {
34 Self::Named(model) => ModelRefWire::Named {
35 model: nonempty_serialize_value::<S::Error>("model", model)?,
36 },
37 Self::BackendModel { backend, model } => ModelRefWire::BackendModel {
38 backend: nonempty_serialize_value::<S::Error>("backend", backend)?,
39 model: nonempty_serialize_value::<S::Error>("model", model)?,
40 },
41 Self::Resolved(_) => {
42 return Err(serde::ser::Error::custom(
43 "resolved ModelRef is process-local and cannot be serialized",
44 ));
45 }
46 };
47 wire.serialize(serializer)
48 }
49}
50
51impl<'de> Deserialize<'de> for ModelRef {
52 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
53 where
54 D: Deserializer<'de>,
55 {
56 match ModelRefWire::deserialize(deserializer)? {
57 ModelRefWire::Named { model } => Ok(Self::Named(
58 nonempty_deserialize_value::<D::Error>("model", &model)?,
59 )),
60 ModelRefWire::BackendModel { backend, model } => Ok(Self::BackendModel {
61 backend: nonempty_deserialize_value::<D::Error>("backend", &backend)?,
62 model: nonempty_deserialize_value::<D::Error>("model", &model)?,
63 }),
64 }
65 }
66}
67
68fn nonempty_serialize_value<E>(field: &str, value: &str) -> Result<String, E>
69where
70 E: serde::ser::Error,
71{
72 if value.trim().is_empty() {
73 return Err(E::custom(format!("ModelRef {field} cannot be empty")));
74 }
75 Ok(value.to_string())
76}
77
78fn nonempty_deserialize_value<E>(field: &str, value: &str) -> Result<String, E>
79where
80 E: serde::de::Error,
81{
82 if value.trim().is_empty() {
83 return Err(E::custom(format!("ModelRef {field} cannot be empty")));
84 }
85 Ok(value.to_string())
86}
87
88impl ModelRef {
89 pub fn named(model: impl Into<String>) -> Self {
90 Self::Named(model.into())
91 }
92
93 pub fn backend(backend: impl Into<String>, model: impl Into<String>) -> Self {
94 Self::BackendModel {
95 backend: backend.into(),
96 model: model.into(),
97 }
98 }
99
100 pub fn resolved(resolved: ResolvedModelConfig) -> Self {
101 Self::Resolved(resolved)
102 }
103
104 pub fn model(&self) -> &str {
105 match self {
106 Self::Named(model) => model,
107 Self::BackendModel { model, .. } => model,
108 Self::Resolved(resolved) => resolved.selected_model.as_str(),
109 }
110 }
111
112 pub fn backend_name(&self) -> Option<&str> {
113 match self {
114 Self::Named(_) => None,
115 Self::BackendModel { backend, .. } => Some(backend),
116 Self::Resolved(resolved) => Some(resolved.backend.as_str()),
117 }
118 }
119}
120
121#[derive(Debug, Error)]
122pub enum ModelError {
123 #[error("model provider has no default backend for named model `{0}`")]
124 MissingDefaultBackend(String),
125 #[error("model backend mismatch: requested `{requested}`, provider `{provider}`")]
126 BackendMismatch { requested: String, provider: String },
127 #[error("{0}")]
128 Config(String),
129}
130
131pub trait ModelProvider: Send + Sync {
132 fn resolve(&self, model: &ModelRef) -> Result<ResolvedModelConfig, ModelError>;
133 fn client(&self, resolved: &ResolvedModelConfig) -> Result<Arc<dyn LlmClient>, ModelError>;
134
135 fn default_settings(&self, _resolved: &ResolvedModelConfig) -> ModelSettings {
136 ModelSettings::default()
137 }
138
139 fn default_model_ref(&self) -> Option<ModelRef> {
140 None
141 }
142}
143
144#[derive(Clone)]
145pub struct ScriptedModelProvider {
146 backend: String,
147 default_model: String,
148 llm: ScriptedLlmClient,
149 context_length: Option<u64>,
150 max_output_tokens: Option<u64>,
151 default_settings: ModelSettings,
152}
153
154impl ScriptedModelProvider {
155 pub fn new(
156 backend: impl Into<String>,
157 default_model: impl Into<String>,
158 responses: Vec<LLMResponse>,
159 ) -> Self {
160 Self::from_steps(
161 backend,
162 default_model,
163 responses.into_iter().map(ScriptStep::response).collect(),
164 )
165 }
166
167 pub fn from_steps(
168 backend: impl Into<String>,
169 default_model: impl Into<String>,
170 steps: Vec<ScriptStep>,
171 ) -> Self {
172 Self {
173 backend: backend.into(),
174 default_model: default_model.into(),
175 llm: ScriptedLlmClient::from_steps(steps),
176 context_length: Some(128_000),
177 max_output_tokens: Some(16_384),
178 default_settings: ModelSettings::default(),
179 }
180 }
181
182 pub fn from_callback(
183 backend: impl Into<String>,
184 default_model: impl Into<String>,
185 callback: impl Fn(&LlmRequest) -> Result<LLMResponse, crate::llm::LlmError>
186 + Send
187 + Sync
188 + 'static,
189 ) -> Self {
190 Self::from_steps(backend, default_model, vec![ScriptStep::callback(callback)])
191 }
192
193 pub fn with_default_settings(mut self, settings: ModelSettings) -> Self {
194 self.default_settings = settings;
195 self
196 }
197
198 pub fn with_token_limits(
199 mut self,
200 context_length: Option<u64>,
201 max_output_tokens: Option<u64>,
202 ) -> Self {
203 self.context_length = context_length;
204 self.max_output_tokens = max_output_tokens;
205 self
206 }
207}
208
209impl ModelProvider for ScriptedModelProvider {
210 fn resolve(&self, model: &ModelRef) -> Result<ResolvedModelConfig, ModelError> {
211 match model {
212 ModelRef::Named(model) => Ok(ResolvedModelConfig::new(
213 self.backend.clone(),
214 model.clone(),
215 model.clone(),
216 model.clone(),
217 Vec::new(),
218 )
219 .with_token_limits(self.context_length, self.max_output_tokens)
220 .with_capabilities(true, true, false)),
221 ModelRef::BackendModel { backend, model } => {
222 if backend != &self.backend {
223 return Err(ModelError::BackendMismatch {
224 requested: backend.clone(),
225 provider: self.backend.clone(),
226 });
227 }
228 Ok(ResolvedModelConfig::new(
229 backend.clone(),
230 model.clone(),
231 model.clone(),
232 model.clone(),
233 Vec::new(),
234 )
235 .with_token_limits(self.context_length, self.max_output_tokens)
236 .with_capabilities(true, true, false))
237 }
238 ModelRef::Resolved(resolved) => Ok(resolved.clone()),
239 }
240 }
241
242 fn client(&self, _resolved: &ResolvedModelConfig) -> Result<Arc<dyn LlmClient>, ModelError> {
243 Ok(Arc::new(self.llm.clone()))
244 }
245
246 fn default_settings(&self, _resolved: &ResolvedModelConfig) -> ModelSettings {
247 self.default_settings.clone()
248 }
249
250 fn default_model_ref(&self) -> Option<ModelRef> {
251 Some(ModelRef::named(self.default_model.clone()))
252 }
253}
254
255impl Default for ScriptedModelProvider {
256 fn default() -> Self {
257 Self::new("scripted", "demo-model", Vec::new())
258 }
259}
260
261#[derive(Debug, Clone)]
262pub struct VvLlmModelProvider {
263 settings_file: PathBuf,
264 default_backend: Option<String>,
265 timeout_seconds: f64,
266}
267
268impl VvLlmModelProvider {
269 pub fn from_settings_file(path: impl Into<PathBuf>) -> Self {
270 Self {
271 settings_file: path.into(),
272 default_backend: None,
273 timeout_seconds: 90.0,
274 }
275 }
276
277 pub fn with_default_backend(mut self, backend: impl Into<String>) -> Self {
278 self.default_backend = Some(backend.into());
279 self
280 }
281
282 pub fn with_timeout_seconds(mut self, timeout_seconds: f64) -> Self {
283 self.timeout_seconds = timeout_seconds.max(1.0);
284 self
285 }
286
287 pub fn settings_file(&self) -> &Path {
288 &self.settings_file
289 }
290}
291
292impl ModelProvider for VvLlmModelProvider {
293 fn resolve(&self, model: &ModelRef) -> Result<ResolvedModelConfig, ModelError> {
294 match model {
295 ModelRef::Named(model) => {
296 let Some(backend) = self.default_backend.as_deref() else {
297 return Err(ModelError::MissingDefaultBackend(model.clone()));
298 };
299 let settings = crate::config::load_llm_settings_from_file(&self.settings_file)
300 .map_err(|error| ModelError::Config(error.to_string()))?;
301 resolve_model_endpoint(&settings, backend, model)
302 .map_err(|error| ModelError::Config(error.to_string()))
303 }
304 ModelRef::BackendModel { backend, model } => {
305 let settings = crate::config::load_llm_settings_from_file(&self.settings_file)
306 .map_err(|error| ModelError::Config(error.to_string()))?;
307 resolve_model_endpoint(&settings, backend, model)
308 .map_err(|error| ModelError::Config(error.to_string()))
309 }
310 ModelRef::Resolved(resolved) => Ok(resolved.clone()),
311 }
312 }
313
314 fn client(&self, resolved: &ResolvedModelConfig) -> Result<Arc<dyn LlmClient>, ModelError> {
315 let (llm, _) = build_vv_llm_from_local_settings(
316 &self.settings_file,
317 &resolved.backend,
318 &resolved.selected_model,
319 self.timeout_seconds,
320 )
321 .map_err(|error| ModelError::Config(error.to_string()))?;
322 let llm: VvLlmClient = llm;
323 Ok(Arc::new(llm) as Arc<dyn LlmClient>)
324 }
325}