diff options
Diffstat (limited to 'src/client/model.rs')
| -rw-r--r-- | src/client/model.rs | 69 |
1 files changed, 24 insertions, 45 deletions
diff --git a/src/client/model.rs b/src/client/model.rs index 3f2cbdd..459d94e 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -3,7 +3,7 @@ use super::message::{Message, MessageContent}; use crate::utils::count_tokens; use anyhow::{bail, Result}; -use serde::{Deserialize, Deserializer}; +use serde::Deserialize; const PER_MESSAGES_TOKENS: usize = 5; const BASIS_TOKENS: usize = 2; @@ -41,10 +41,10 @@ impl Model { .iter() .map(|v| { Model::new(client_name, &v.name) - .set_capabilities(v.capabilities) .set_max_input_tokens(v.max_input_tokens) .set_max_output_tokens(v.max_output_tokens) - .set_extra_fields(v.extra_fields.clone()) + .set_supports_vision(v.supports_vision) + .set_extra_fields(&v.extra_fields) }) .collect() } @@ -84,19 +84,6 @@ impl Model { format!("{}:{}", self.client_name, self.name) } - pub fn set_capabilities(mut self, capabilities: ModelCapabilities) -> Self { - self.capabilities = capabilities; - self - } - - pub fn set_extra_fields( - mut self, - extra_fields: Option<serde_json::Map<String, serde_json::Value>>, - ) -> Self { - self.extra_fields = extra_fields; - self - } - pub fn set_max_input_tokens(mut self, max_input_tokens: Option<usize>) -> Self { match max_input_tokens { None | Some(0) => self.max_input_tokens = None, @@ -113,6 +100,23 @@ impl Model { self } + pub fn set_supports_vision(mut self, supports_vision: bool) -> Self { + if supports_vision { + self.capabilities |= ModelCapabilities::Vision; + } else { + self.capabilities &= !ModelCapabilities::Vision; + } + self + } + + pub fn set_extra_fields( + mut self, + extra_fields: &Option<serde_json::Map<String, serde_json::Value>>, + ) -> Self { + self.extra_fields = extra_fields.clone(); + self + } + pub fn messages_tokens(&self, messages: &[Message]) -> usize { messages .iter() @@ -174,10 +178,11 @@ pub struct ModelConfig { pub name: String, pub max_input_tokens: Option<usize>, pub max_output_tokens: Option<isize>, + pub input_price: Option<f64>, + pub output_price: Option<f64>, + #[serde(default)] + pub supports_vision: bool, pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>, - #[serde(deserialize_with = "deserialize_capabilities")] - #[serde(default = "default_capabilities")] - pub capabilities: ModelCapabilities, } bitflags::bitflags! { @@ -187,29 +192,3 @@ bitflags::bitflags! { const Vision = 0b00000010; } } - -impl From<&str> for ModelCapabilities { - fn from(value: &str) -> Self { - let value = if value.is_empty() { "text" } else { value }; - let mut output = ModelCapabilities::empty(); - if value.contains("text") { - output |= ModelCapabilities::Text; - } - if value.contains("vision") { - output |= ModelCapabilities::Vision; - } - output - } -} - -fn deserialize_capabilities<'de, D>(deserializer: D) -> Result<ModelCapabilities, D::Error> -where - D: Deserializer<'de>, -{ - let value: String = Deserialize::deserialize(deserializer)?; - Ok(value.as_str().into()) -} - -fn default_capabilities() -> ModelCapabilities { - ModelCapabilities::Text -} |
