Skip to main content

starweaver_model/wrappers/
profile_override.rs

1//! Capability-profile override model wrapper.
2
3use async_trait::async_trait;
4
5use super::DynModelAdapter;
6use crate::{
7    adapter::{
8        ModelAdapter, ModelError, ModelRequestContext, ModelRequestParameters,
9        ModelResponseEventStream, ModelRunSession,
10    },
11    message::{ModelMessage, ModelResponse},
12    profile::ModelProfile,
13    settings::ModelSettings,
14    stream::ModelResponseStreamEvent,
15};
16
17/// Model wrapper that overlays a capability profile and optional default settings.
18pub struct ProfileOverrideModel {
19    inner: DynModelAdapter,
20    model_name: String,
21    provider_name: Option<String>,
22    profile: ModelProfile,
23    default_settings: Option<ModelSettings>,
24}
25
26impl ProfileOverrideModel {
27    /// Create a wrapper with a replacement profile.
28    #[must_use]
29    pub fn new(inner: DynModelAdapter, profile: ModelProfile) -> Self {
30        Self {
31            model_name: inner.model_name().to_string(),
32            provider_name: inner.provider_name().map(str::to_string),
33            default_settings: inner.default_settings().cloned(),
34            inner,
35            profile,
36        }
37    }
38
39    /// Override the exposed model name.
40    #[must_use]
41    pub fn with_model_name(mut self, model_name: impl Into<String>) -> Self {
42        self.model_name = model_name.into();
43        self
44    }
45
46    /// Override the exposed provider name.
47    #[must_use]
48    pub fn with_provider_name(mut self, provider_name: impl Into<Option<String>>) -> Self {
49        self.provider_name = provider_name.into();
50        self
51    }
52
53    /// Override wrapper default settings.
54    #[must_use]
55    pub fn with_default_settings(mut self, settings: ModelSettings) -> Self {
56        self.default_settings = Some(settings);
57        self
58    }
59}
60
61#[async_trait]
62impl ModelAdapter for ProfileOverrideModel {
63    fn model_name(&self) -> &str {
64        &self.model_name
65    }
66
67    fn provider_name(&self) -> Option<&str> {
68        self.provider_name.as_deref()
69    }
70
71    fn profile(&self) -> &ModelProfile {
72        &self.profile
73    }
74
75    fn default_settings(&self) -> Option<&ModelSettings> {
76        self.default_settings.as_ref()
77    }
78
79    fn start_run_session(&self) -> Box<dyn ModelRunSession + '_> {
80        self.inner.start_run_session()
81    }
82
83    async fn request(
84        &self,
85        messages: Vec<ModelMessage>,
86        settings: Option<ModelSettings>,
87        params: ModelRequestParameters,
88        context: ModelRequestContext,
89    ) -> Result<ModelResponse, ModelError> {
90        self.inner
91            .request(messages, settings, params, context)
92            .await
93    }
94
95    async fn request_stream(
96        &self,
97        messages: Vec<ModelMessage>,
98        settings: Option<ModelSettings>,
99        params: ModelRequestParameters,
100        context: ModelRequestContext,
101    ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
102        self.inner
103            .request_stream(messages, settings, params, context)
104            .await
105    }
106
107    async fn request_stream_incremental(
108        &self,
109        messages: Vec<ModelMessage>,
110        settings: Option<ModelSettings>,
111        params: ModelRequestParameters,
112        context: ModelRequestContext,
113    ) -> Result<ModelResponseEventStream, ModelError> {
114        self.inner
115            .request_stream_incremental(messages, settings, params, context)
116            .await
117    }
118
119    async fn count_tokens(
120        &self,
121        messages: &[ModelMessage],
122        settings: Option<&ModelSettings>,
123        params: &ModelRequestParameters,
124    ) -> Result<starweaver_usage::Usage, ModelError> {
125        self.inner.count_tokens(messages, settings, params).await
126    }
127}