Skip to main content

codewhale_workflow/
model_policy.rs

1use std::collections::BTreeMap;
2
3use serde::de::DeserializeOwned;
4use serde::{Deserialize, Serialize};
5use thiserror::Error;
6
7use crate::{AgentType, ModelPolicy, WorkflowUsage};
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
10#[serde(rename_all = "snake_case")]
11pub enum ModelRole {
12    Planner,
13    LeafReasoner,
14    Implementer,
15    Reviewer,
16    Teacher,
17    Student,
18    JsonExtractor,
19}
20
21impl From<AgentType> for ModelRole {
22    fn from(agent_type: AgentType) -> Self {
23        match agent_type {
24            AgentType::General | AgentType::Explore => Self::LeafReasoner,
25            AgentType::Plan => Self::Planner,
26            AgentType::Review | AgentType::Verifier => Self::Reviewer,
27            AgentType::Implementer => Self::Implementer,
28        }
29    }
30}
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
33pub struct ModelCapabilities {
34    #[serde(default)]
35    pub tool_calls: bool,
36    #[serde(default)]
37    pub json_mode: bool,
38    #[serde(default)]
39    pub prompt_cache: bool,
40    #[serde(default)]
41    pub large_context: bool,
42    #[serde(default)]
43    pub streaming: bool,
44}
45
46impl ModelCapabilities {
47    #[must_use]
48    pub fn satisfies(self, required: Self) -> bool {
49        (!required.tool_calls || self.tool_calls)
50            && (!required.json_mode || self.json_mode)
51            && (!required.prompt_cache || self.prompt_cache)
52            && (!required.large_context || self.large_context)
53            && (!required.streaming || self.streaming)
54    }
55}
56
57#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
58pub struct ProviderModel {
59    pub provider: String,
60    pub model: String,
61    #[serde(default)]
62    pub capabilities: ModelCapabilities,
63}
64
65#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
66pub struct ResolvedModel {
67    pub role: ModelRole,
68    pub provider: String,
69    pub model: String,
70    pub capabilities: ModelCapabilities,
71    pub source: ModelSelectionSource,
72}
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
75#[serde(rename_all = "snake_case")]
76pub enum ModelSelectionSource {
77    Primary,
78    Fallback,
79    RoleDefault,
80}
81
82#[derive(Debug, Clone, Default)]
83pub struct ProviderRegistry {
84    models: BTreeMap<String, ProviderModel>,
85    role_policies: BTreeMap<ModelRole, ModelPolicy>,
86}
87
88impl ProviderRegistry {
89    pub fn new() -> Self {
90        Self::default()
91    }
92
93    pub fn with_model(mut self, model: ProviderModel) -> Self {
94        self.insert_model(model);
95        self
96    }
97
98    pub fn with_role_policy(mut self, role: ModelRole, policy: ModelPolicy) -> Self {
99        self.role_policies.insert(role, policy);
100        self
101    }
102
103    pub fn insert_model(&mut self, model: ProviderModel) {
104        self.models
105            .insert(model_key(&model.provider, &model.model), model);
106    }
107
108    pub fn resolve_role(
109        &self,
110        role: ModelRole,
111        policy: Option<&ModelPolicy>,
112        required: ModelCapabilities,
113    ) -> Result<ResolvedModel, ModelPolicyError> {
114        let policy = match policy {
115            Some(policy) => (policy, ModelSelectionSource::Primary),
116            None => (
117                self.role_policies
118                    .get(&role)
119                    .ok_or(ModelPolicyError::MissingPolicy { role })?,
120                ModelSelectionSource::RoleDefault,
121            ),
122        };
123        self.resolve_policy(role, policy.0, policy.1, required)
124    }
125
126    fn resolve_policy(
127        &self,
128        role: ModelRole,
129        policy: &ModelPolicy,
130        primary_source: ModelSelectionSource,
131        required: ModelCapabilities,
132    ) -> Result<ResolvedModel, ModelPolicyError> {
133        let candidates = model_candidates(policy)?;
134        let mut rejected = Vec::new();
135        for (index, candidate) in candidates.iter().enumerate() {
136            let source = if index == 0 {
137                primary_source
138            } else {
139                ModelSelectionSource::Fallback
140            };
141            let Some(model) = self
142                .models
143                .get(&model_key(&candidate.provider, &candidate.model))
144            else {
145                rejected.push(format!(
146                    "{}/{}: unknown",
147                    candidate.provider, candidate.model
148                ));
149                continue;
150            };
151            if model.capabilities.satisfies(required) {
152                return Ok(ResolvedModel {
153                    role,
154                    provider: model.provider.clone(),
155                    model: model.model.clone(),
156                    capabilities: model.capabilities,
157                    source,
158                });
159            }
160            rejected.push(format!(
161                "{}/{}: missing required capabilities",
162                model.provider, model.model
163            ));
164        }
165        Err(ModelPolicyError::NoCapableModel { role, rejected })
166    }
167}
168
169#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
170pub struct CompletionRequest {
171    pub role: ModelRole,
172    pub prompt: String,
173    #[serde(default)]
174    pub require_json: bool,
175    #[serde(default)]
176    pub model_policy: ModelPolicy,
177}
178
179#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
180pub struct CompletionResponse {
181    pub text: String,
182    #[serde(default)]
183    pub usage: WorkflowUsage,
184}
185
186pub trait ModelProvider {
187    fn provider(&self) -> &str;
188    fn model(&self) -> &str;
189    fn capabilities(&self) -> ModelCapabilities;
190    fn complete(
191        &self,
192        request: &CompletionRequest,
193    ) -> Result<CompletionResponse, ModelProviderError>;
194}
195
196#[derive(Debug, Clone)]
197pub struct MockModelProvider {
198    provider: String,
199    model: String,
200    capabilities: ModelCapabilities,
201    response: CompletionResponse,
202}
203
204impl MockModelProvider {
205    pub fn new(
206        provider: impl Into<String>,
207        model: impl Into<String>,
208        capabilities: ModelCapabilities,
209        response: impl Into<String>,
210    ) -> Self {
211        Self {
212            provider: provider.into(),
213            model: model.into(),
214            capabilities,
215            response: CompletionResponse {
216                text: response.into(),
217                usage: WorkflowUsage::default(),
218            },
219        }
220    }
221}
222
223impl ModelProvider for MockModelProvider {
224    fn provider(&self) -> &str {
225        &self.provider
226    }
227
228    fn model(&self) -> &str {
229        &self.model
230    }
231
232    fn capabilities(&self) -> ModelCapabilities {
233        self.capabilities
234    }
235
236    fn complete(
237        &self,
238        _request: &CompletionRequest,
239    ) -> Result<CompletionResponse, ModelProviderError> {
240        Ok(self.response.clone())
241    }
242}
243
244#[derive(Debug, Clone, PartialEq, Eq, Error)]
245pub enum ModelPolicyError {
246    #[error("no model policy configured for role `{role:?}`")]
247    MissingPolicy { role: ModelRole },
248    #[error("model policy must include a model for role resolution")]
249    MissingModel,
250    #[error("fallback model `{model}` requires a provider when the primary policy has none")]
251    MissingFallbackProvider { model: String },
252    #[error("no configured model satisfies role `{role:?}` requirements: {rejected:?}")]
253    NoCapableModel {
254        role: ModelRole,
255        rejected: Vec<String>,
256    },
257}
258
259#[derive(Debug, Clone, PartialEq, Eq, Error)]
260pub enum ModelProviderError {
261    #[error("model provider `{provider}/{model}` failed: {reason}")]
262    Failed {
263        provider: String,
264        model: String,
265        reason: String,
266    },
267}
268
269#[derive(Debug, Clone, PartialEq, Eq, Error)]
270pub enum JsonRepairError {
271    #[error("json parse failed before and after one repair pass: {reason}")]
272    Parse { reason: String },
273}
274
275pub fn parse_json_with_repair<T: DeserializeOwned>(raw: &str) -> Result<T, JsonRepairError> {
276    match serde_json::from_str(raw) {
277        Ok(parsed) => Ok(parsed),
278        Err(first) => {
279            let repaired = repair_json_text_once(raw);
280            serde_json::from_str(&repaired).map_err(|second| JsonRepairError::Parse {
281                reason: format!("{first}; repair failed: {second}"),
282            })
283        }
284    }
285}
286
287pub fn repair_json_text_once(raw: &str) -> String {
288    let trimmed = raw.trim();
289    let without_fence = trimmed
290        .strip_prefix("```json")
291        .or_else(|| trimmed.strip_prefix("```"))
292        .and_then(|value| value.strip_suffix("```"))
293        .map(str::trim)
294        .unwrap_or(trimmed);
295    let object = slice_json_payload(without_fence, '{', '}');
296    let array = slice_json_payload(without_fence, '[', ']');
297    object.or(array).unwrap_or(without_fence).to_string()
298}
299
300#[derive(Debug, Clone, PartialEq, Eq)]
301struct ModelCandidate {
302    provider: String,
303    model: String,
304}
305
306fn model_candidates(policy: &ModelPolicy) -> Result<Vec<ModelCandidate>, ModelPolicyError> {
307    let mut candidates = Vec::new();
308    let Some(primary_model) = policy.model.as_ref() else {
309        return Err(ModelPolicyError::MissingModel);
310    };
311    candidates.push(candidate_from_model(
312        policy.provider.as_deref(),
313        primary_model,
314    )?);
315    for fallback in &policy.fallback_models {
316        candidates.push(candidate_from_model(policy.provider.as_deref(), fallback)?);
317    }
318    Ok(candidates)
319}
320
321fn candidate_from_model(
322    default_provider: Option<&str>,
323    model: &str,
324) -> Result<ModelCandidate, ModelPolicyError> {
325    if let Some((provider, model)) = model.split_once('/') {
326        return Ok(ModelCandidate {
327            provider: provider.to_string(),
328            model: model.to_string(),
329        });
330    }
331    let Some(provider) = default_provider else {
332        return Err(ModelPolicyError::MissingFallbackProvider {
333            model: model.to_string(),
334        });
335    };
336    Ok(ModelCandidate {
337        provider: provider.to_string(),
338        model: model.to_string(),
339    })
340}
341
342fn model_key(provider: &str, model: &str) -> String {
343    format!("{provider}/{model}")
344}
345
346fn slice_json_payload(raw: &str, open: char, close: char) -> Option<&str> {
347    let start = raw.find(open)?;
348    let end = raw.rfind(close)?;
349    (end >= start).then_some(&raw[start..=end])
350}
351
352#[cfg(test)]
353mod tests {
354    use super::*;
355
356    fn model(provider: &str, model: &str, capabilities: ModelCapabilities) -> ProviderModel {
357        ProviderModel {
358            provider: provider.to_string(),
359            model: model.to_string(),
360            capabilities,
361        }
362    }
363
364    #[test]
365    fn provider_capability_fallback() {
366        let registry = ProviderRegistry::new()
367            .with_model(model("mock", "plain", ModelCapabilities::default()))
368            .with_model(model(
369                "mock",
370                "json",
371                ModelCapabilities {
372                    json_mode: true,
373                    ..ModelCapabilities::default()
374                },
375            ));
376        let policy = ModelPolicy {
377            provider: Some("mock".to_string()),
378            model: Some("plain".to_string()),
379            fallback_models: vec!["json".to_string()],
380        };
381
382        let resolved = registry
383            .resolve_role(
384                ModelRole::JsonExtractor,
385                Some(&policy),
386                ModelCapabilities {
387                    json_mode: true,
388                    ..ModelCapabilities::default()
389                },
390            )
391            .expect("fallback json model should satisfy the role");
392
393        assert_eq!(resolved.model, "json");
394        assert_eq!(resolved.source, ModelSelectionSource::Fallback);
395    }
396
397    #[test]
398    fn role_default_policy_resolves_model() {
399        let registry = ProviderRegistry::new()
400            .with_model(model(
401                "mock",
402                "planner",
403                ModelCapabilities {
404                    large_context: true,
405                    ..ModelCapabilities::default()
406                },
407            ))
408            .with_role_policy(
409                ModelRole::Planner,
410                ModelPolicy {
411                    provider: Some("mock".to_string()),
412                    model: Some("planner".to_string()),
413                    fallback_models: Vec::new(),
414                },
415            );
416
417        let resolved = registry
418            .resolve_role(
419                ModelRole::Planner,
420                None,
421                ModelCapabilities {
422                    large_context: true,
423                    ..ModelCapabilities::default()
424                },
425            )
426            .expect("role default should resolve");
427
428        assert_eq!(resolved.role, ModelRole::Planner);
429        assert_eq!(resolved.source, ModelSelectionSource::RoleDefault);
430    }
431
432    #[test]
433    fn agent_type_maps_to_model_role() {
434        assert_eq!(ModelRole::from(AgentType::Plan), ModelRole::Planner);
435        assert_eq!(
436            ModelRole::from(AgentType::Implementer),
437            ModelRole::Implementer
438        );
439        assert_eq!(ModelRole::from(AgentType::Verifier), ModelRole::Reviewer);
440    }
441
442    #[test]
443    fn json_repair_fallback() {
444        #[derive(Debug, Deserialize, PartialEq, Eq)]
445        struct Payload {
446            answer: String,
447        }
448
449        let parsed: Payload = parse_json_with_repair(
450            r#"Here is the JSON:
451```json
452{"answer":"ok"}
453```
454"#,
455        )
456        .expect("repair should extract fenced JSON");
457
458        assert_eq!(
459            parsed,
460            Payload {
461                answer: "ok".to_string()
462            }
463        );
464    }
465
466    #[test]
467    fn json_repair_fallback_fails_closed() {
468        let err = parse_json_with_repair::<serde_json::Value>("not json")
469            .expect_err("non-json text should fail closed");
470
471        assert!(matches!(err, JsonRepairError::Parse { .. }));
472    }
473
474    #[test]
475    fn mock_provider_returns_configured_response() {
476        let provider = MockModelProvider::new(
477            "mock",
478            "fast",
479            ModelCapabilities::default(),
480            "mock response",
481        );
482        let request = CompletionRequest {
483            role: ModelRole::LeafReasoner,
484            prompt: "say something".to_string(),
485            require_json: false,
486            model_policy: ModelPolicy::default(),
487        };
488
489        let response = provider.complete(&request).expect("mock should respond");
490
491        assert_eq!(provider.provider(), "mock");
492        assert_eq!(provider.model(), "fast");
493        assert_eq!(response.text, "mock response");
494    }
495}