Skip to main content

relay_knowledge/model_provider/
profiles.rs

1//! Owns persisted model profile workflows and runtime-profile resolution.
2
3use std::collections::BTreeMap;
4
5use tokio::fs;
6
7use super::{
8    ModelProfileRuntimeSummary, ModelProfileSaveRequest, ModelProfilesResponse,
9    ModelProviderConfigService, ModelProviderError,
10    persistence::write_json,
11    profile::{DEFAULT_PROFILE_NAME, StoredModelProfile, StoredProfileFile},
12    profile_config::profile_response,
13    profile_config::runtime_profile_merge_base,
14    profile_config::validate_profile_name,
15};
16use crate::retrieval::ReadModelBackendConfig;
17
18impl ModelProviderConfigService {
19    pub async fn profiles(
20        &self,
21        retrieval: &ReadModelBackendConfig,
22    ) -> Result<ModelProfilesResponse, ModelProviderError> {
23        let file = self.load_profile_file().await?;
24        Ok(profile_response(file, retrieval))
25    }
26
27    pub async fn profile_summary(
28        &self,
29        retrieval: &ReadModelBackendConfig,
30    ) -> ModelProfileRuntimeSummary {
31        match self.profiles(retrieval).await {
32            Ok(response) => ModelProfileRuntimeSummary {
33                loaded: response.loaded,
34                profile_count: response.profiles.len(),
35                default_profile: response.default_profile,
36                error: response.error,
37            },
38            Err(error) => ModelProfileRuntimeSummary {
39                loaded: false,
40                profile_count: 0,
41                default_profile: None,
42                error: Some(error.to_string()),
43            },
44        }
45    }
46
47    pub async fn save_profile(
48        &self,
49        name: &str,
50        request: ModelProfileSaveRequest,
51        retrieval: &ReadModelBackendConfig,
52    ) -> Result<ModelProfilesResponse, ModelProviderError> {
53        let name = validate_profile_name(name)?;
54        let mut file = self
55            .load_profile_file()
56            .await?
57            .unwrap_or_else(|| StoredProfileFile {
58                default_profile: None,
59                profiles: BTreeMap::new(),
60            });
61        let runtime_profile = runtime_profile_merge_base(&file, &name, retrieval);
62        let existing = file.profiles.get(&name).or(runtime_profile.as_ref());
63        let is_default = request.is_default || file.default_profile.is_none();
64        let stored = StoredModelProfile::from_save_request(request, existing)?;
65        file.profiles.insert(name.clone(), stored);
66        if is_default {
67            file.default_profile = Some(name);
68            for (profile_name, profile) in &mut file.profiles {
69                profile.is_default = file.default_profile.as_ref() == Some(profile_name);
70            }
71        }
72        self.write_profile_file(&file).await?;
73        Ok(profile_response(Some(file), retrieval))
74    }
75
76    pub async fn delete_profile(
77        &self,
78        name: &str,
79        retrieval: &ReadModelBackendConfig,
80    ) -> Result<ModelProfilesResponse, ModelProviderError> {
81        let name = validate_profile_name(name)?;
82        let mut file = self
83            .load_profile_file()
84            .await?
85            .unwrap_or_else(|| StoredProfileFile {
86                default_profile: None,
87                profiles: BTreeMap::new(),
88            });
89        file.profiles.remove(&name);
90        if file.default_profile.as_deref() == Some(&name) {
91            file.default_profile = file.profiles.keys().next().cloned();
92        }
93        for (profile_name, profile) in &mut file.profiles {
94            profile.is_default = file.default_profile.as_ref() == Some(profile_name);
95        }
96        self.write_profile_file(&file).await?;
97        Ok(profile_response(Some(file), retrieval))
98    }
99
100    pub(super) async fn resolve_probe_profile(
101        &self,
102        retrieval: &ReadModelBackendConfig,
103        profile_name: Option<String>,
104        override_config: Option<ModelProfileSaveRequest>,
105    ) -> Result<StoredModelProfile, ModelProviderError> {
106        match (profile_name, override_config) {
107            (Some(name), Some(request)) => {
108                let base = self.resolve_profile_by_name(retrieval, &name).await?;
109                StoredModelProfile::from_save_request(request, Some(&base))
110            }
111            (Some(name), None) => self.resolve_profile_by_name(retrieval, &name).await,
112            (None, Some(request)) => {
113                let base = match self.resolve_default_profile(retrieval).await {
114                    Ok(profile) => Some(profile),
115                    Err(ModelProviderError::InvalidInput(message))
116                        if message == "no model profile is configured" =>
117                    {
118                        None
119                    }
120                    Err(error) => return Err(error),
121                };
122                StoredModelProfile::from_save_request(request, base.as_ref())
123            }
124            (None, None) => self.resolve_default_profile(retrieval).await,
125        }
126    }
127
128    async fn resolve_default_profile(
129        &self,
130        retrieval: &ReadModelBackendConfig,
131    ) -> Result<StoredModelProfile, ModelProviderError> {
132        let file = self.load_profile_file().await?;
133        let response = profile_response(file.clone(), retrieval);
134        let Some(default_name) = response.default_profile else {
135            return Err(ModelProviderError::InvalidInput(
136                "no model profile is configured".to_owned(),
137            ));
138        };
139        self.resolve_profile_by_name(retrieval, &default_name).await
140    }
141
142    async fn resolve_profile_by_name(
143        &self,
144        retrieval: &ReadModelBackendConfig,
145        name: &str,
146    ) -> Result<StoredModelProfile, ModelProviderError> {
147        let name = validate_profile_name(name)?;
148        if let Some(file) = self.load_profile_file().await? {
149            if let Some(profile) = file.profiles.get(&name) {
150                return Ok(profile.clone());
151            }
152        }
153        if name == DEFAULT_PROFILE_NAME {
154            if let Some(profile) = StoredModelProfile::from_runtime(retrieval) {
155                return Ok(profile);
156            }
157        }
158        Err(ModelProviderError::InvalidInput(format!(
159            "model profile '{name}' was not found"
160        )))
161    }
162
163    async fn load_profile_file(&self) -> Result<Option<StoredProfileFile>, ModelProviderError> {
164        match fs::read_to_string(self.paths.model_profiles_file()).await {
165            Ok(raw) => serde_json::from_str(&raw)
166                .map(Some)
167                .map_err(ModelProviderError::from),
168            Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(None),
169            Err(error) => Err(ModelProviderError::from(error)),
170        }
171    }
172
173    async fn write_profile_file(&self, file: &StoredProfileFile) -> Result<(), ModelProviderError> {
174        write_json(self.paths.model_profiles_file(), file).await
175    }
176}
177
178#[cfg(test)]
179#[path = "profiles_tests.rs"]
180mod tests;