Skip to main content

vv_agent/
model.rs

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}