starweaver_model/wrappers/
profile_override.rs1use 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
17pub 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 #[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 #[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 #[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 #[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}