From fe35cfd9419302f01baf9672493c0b0a4b41d889 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 13 Jan 2024 19:52:07 +0800 Subject: 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 --- src/client/model.rs | 51 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 51 insertions(+) (limited to 'src/client/model.rs') 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, 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) -> 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, + #[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 +where + D: Deserializer<'de>, +{ + let value: String = Deserialize::deserialize(deserializer)?; + Ok(value.as_str().into()) +} + +fn default_capabilities() -> ModelCapabilities { + ModelCapabilities::Text +} -- cgit v1.2.3