summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-01-13 19:52:07 +0800
committerGitHub <noreply@github.com>2024-01-13 19:52:07 +0800
commitfe35cfd9419302f01baf9672493c0b0a4b41d889 (patch)
tree94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/model.rs
parent4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff)
downloadaichat-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.rs51
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
+}