diff options
| author | sigoden <sigoden@gmail.com> | 2024-01-13 19:52:07 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-01-13 19:52:07 +0800 |
| commit | fe35cfd9419302f01baf9672493c0b0a4b41d889 (patch) | |
| tree | 94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/model.rs | |
| parent | 4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff) | |
| download | aichat-fe35cfd9419302f01baf9672493c0b0a4b41d889.tar.gz | |
feat: supports model capabilities (#297)
1. automatically switch to the model that has the necessary capabilities.
2. throw an error if the client does not have a model with the necessary capabilities
Diffstat (limited to 'src/client/model.rs')
| -rw-r--r-- | src/client/model.rs | 51 |
1 files changed, 51 insertions, 0 deletions
diff --git a/src/client/model.rs b/src/client/model.rs index 130489d..88d3294 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -3,6 +3,7 @@ use super::message::{Message, MessageContent}; use crate::utils::count_tokens; use anyhow::{bail, Result}; +use serde::{Deserialize, Deserializer}; pub type TokensCountFactors = (usize, usize); // (per-messages, bias) @@ -12,6 +13,7 @@ pub struct Model { pub name: String, pub max_tokens: Option<usize>, pub tokens_count_factors: TokensCountFactors, + pub capabilities: ModelCapabilities, } impl Default for Model { @@ -27,6 +29,7 @@ impl Model { name: name.into(), max_tokens: None, tokens_count_factors: Default::default(), + capabilities: ModelCapabilities::Text, } } @@ -65,6 +68,11 @@ impl Model { format!("{}:{}", self.client_name, self.name) } + pub fn set_capabilities(mut self, capabilities: ModelCapabilities) -> Self { + self.capabilities = capabilities; + self + } + pub fn set_max_tokens(mut self, max_tokens: Option<usize>) -> Self { match max_tokens { None | Some(0) => self.max_tokens = None, @@ -115,3 +123,46 @@ impl Model { Ok(()) } } + +#[derive(Debug, Clone, Deserialize)] +pub struct ModelConfig { + pub name: String, + pub max_tokens: Option<usize>, + #[serde(deserialize_with = "deserialize_capabilities")] + #[serde(default = "default_capabilities")] + pub capabilities: ModelCapabilities, +} + +bitflags::bitflags! { + #[derive(Debug, Clone, Copy, PartialEq)] + pub struct ModelCapabilities: u32 { + const Text = 0b00000001; + 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 +} |
