From 4db9b309803796bc5f996d0b3713344eb44207ec Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 25 Apr 2024 10:15:54 +0800 Subject: refactor: rewrite list models of all clients (#436) --- src/client/claude.rs | 38 +++++++++--------- src/client/cohere.rs | 34 ++++++++-------- src/client/common.rs | 34 +++++----------- src/client/ernie.rs | 106 +++++++++++++++---------------------------------- src/client/gemini.rs | 17 ++++---- src/client/message.rs | 21 ++++++++++ src/client/mistral.rs | 17 ++++---- src/client/model.rs | 11 ----- src/client/ollama.rs | 13 +++--- src/client/openai.rs | 42 ++++++++++---------- src/client/qianwen.rs | 38 +++++++++--------- src/client/vertexai.rs | 31 +++++++-------- 12 files changed, 176 insertions(+), 226 deletions(-) (limited to 'src/client') diff --git a/src/client/claude.rs b/src/client/claude.rs index 054731c..68ab509 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -15,13 +15,6 @@ use serde_json::{json, Value}; const API_BASE: &str = "https://api.anthropic.com/v1/messages"; -const MODELS: [(&str, usize, &str); 3] = [ - // https://docs.anthropic.com/claude/docs/models-overview - ("claude-3-opus-20240229", 200000, "text,vision"), - ("claude-3-sonnet-20240229", 200000, "text,vision"), - ("claude-3-haiku-20240307", 200000, "text,vision"), -]; - #[derive(Debug, Clone, Deserialize)] pub struct ClaudeConfig { pub name: Option, @@ -52,7 +45,16 @@ impl Client for ClaudeClient { } impl ClaudeClient { - list_models_fn!(ClaudeConfig, &MODELS); + list_models_fn!( + ClaudeConfig, + [ + // https://docs.anthropic.com/claude/docs/models-overview + ("claude-3-opus-20240229", "text,vision", 200000, 4096), + ("claude-3-sonnet-20240229", "text,vision", 200000, 4096), + ("claude-3-haiku-20240307", "text,vision", 200000, 4096), + ] + ); + config_get_fn!(api_key, get_api_key); pub const PROMPTS: [PromptType<'static>; 1] = @@ -61,7 +63,7 @@ impl ClaudeClient { fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { let api_key = self.get_api_key().ok(); - let body = build_body(data, &self.model)?; + let body = claude_build_body(data, &self.model)?; let url = API_BASE; @@ -136,7 +138,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand Ok(()) } -fn build_body(data: SendData, model: &Model) -> Result { +pub fn claude_build_body(data: SendData, model: &Model) -> Result { let SendData { mut messages, temperature, @@ -191,18 +193,16 @@ fn build_body(data: SendData, model: &Model) -> Result { ); } - let max_tokens = model.max_output_tokens.unwrap_or(4096); - let mut body = json!({ "model": &model.name, - "max_tokens": max_tokens, "messages": messages, }); - - if let Some(system) = system_message { - body["system"] = system.into(); + if let Some(v) = system_message { + body["system"] = v.into(); + } + if let Some(v) = model.max_output_tokens { + body["max_tokens"] = v.into(); } - if let Some(v) = temperature { body["temperature"] = v.into(); } @@ -218,8 +218,8 @@ fn build_body(data: SendData, model: &Model) -> Result { fn catch_error(data: &Value, status: u16) -> Result<()> { debug!("Invalid response, status: {status}, data: {data}"); if let Some(error) = data["error"].as_object() { - if let (Some(type_), Some(message)) = (error["type"].as_str(), error["message"].as_str()) { - bail!("{message} (type: {type_})"); + if let (Some(typ), Some(message)) = (error["type"].as_str(), error["message"].as_str()) { + bail!("{message} (type: {typ})"); } } bail!("Invalid response, status: {status}, data: {data}"); diff --git a/src/client/cohere.rs b/src/client/cohere.rs index cfab0fa..828b799 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -13,12 +13,6 @@ use serde_json::{json, Value}; const API_URL: &str = "https://api.cohere.ai/v1/chat"; -const MODELS: [(&str, usize, &str); 2] = [ - // https://docs.cohere.com/docs/command-r - ("command-r", 128000, "text"), - ("command-r-plus", 128000, "text"), -]; - #[derive(Debug, Clone, Deserialize, Default)] pub struct CohereConfig { pub name: Option, @@ -49,7 +43,14 @@ impl Client for CohereClient { } impl CohereClient { - list_models_fn!(CohereConfig, &MODELS); + list_models_fn!( + CohereConfig, + [ + // https://docs.cohere.com/docs/command-r + ("command-r", "text", 128000), + ("command-r-plus", "text", 128000), + ] + ); config_get_fn!(api_key, get_api_key); pub const PROMPTS: [PromptType<'static>; 1] = @@ -159,23 +160,22 @@ fn build_body(data: SendData, model: &Model) -> Result { "message": message, }); - if let Some(preamble) = system_message { - body["preamble"] = preamble.into(); - } - - if let Some(max_tokens) = model.max_output_tokens { - body["max_tokens"] = max_tokens.into(); + if let Some(v) = system_message { + body["preamble"] = v.into(); } if !messages.is_empty() { body["chat_history"] = messages.into(); } - if let Some(temperature) = temperature { - body["temperature"] = temperature.into(); + if let Some(v) = model.max_output_tokens { + body["max_tokens"] = v.into(); + } + if let Some(v) = temperature { + body["temperature"] = v.into(); } - if let Some(top_p) = top_p { - body["p"] = top_p.into(); + if let Some(v) = top_p { + body["p"] = v.into(); } if stream { body["stream"] = true.into(); diff --git a/src/client/common.rs b/src/client/common.rs index 593140d..e003e7e 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,4 +1,4 @@ -use super::{openai::OpenAIConfig, ClientConfig, Message, MessageContent, Model, ReplyHandler}; +use super::{openai::OpenAIConfig, ClientConfig, Message, Model, ReplyHandler}; use crate::{ config::{GlobalConfig, Input}, @@ -205,11 +205,18 @@ macro_rules! list_models_fn { Model::from_config(client_name, &local_config.models) } }; - ($config:ident, $models:expr) => { + ($config:ident, [$(($name:literal, $capabilities:literal, $max_input_tokens:literal $(, $max_output_tokens:literal)? )),+$(,)?]) => { pub fn list_models(local_config: &$config) -> Vec { let client_name = Self::name(local_config); if local_config.models.is_empty() { - Model::from_static(client_name, $models) + vec![ + $( + Model::new(client_name, $name) + .set_capabilities($capabilities.into()) + .set_max_input_tokens(Some($max_input_tokens)) + $(.set_max_output_tokens(Some($max_output_tokens)))? + ),+ + ] } else { Model::from_config(client_name, &local_config.models) } @@ -402,27 +409,6 @@ where Ok(()) } -pub fn patch_system_message(messages: &mut Vec) { - if messages[0].role.is_system() { - let system_message = messages.remove(0); - if let (Some(message), MessageContent::Text(system_text)) = - (messages.get_mut(0), system_message.content) - { - if let MessageContent::Text(text) = message.content.clone() { - message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text)) - } - } - } -} - -pub fn extract_sytem_message(messages: &mut Vec) -> Option { - if messages[0].role.is_system() { - let system_message = messages.remove(0); - return Some(system_message.content.to_text()); - } - None -} - pub async fn json_stream(mut stream: S, mut handle: F) -> Result<()> where S: Stream> + Unpin, diff --git a/src/client/ernie.rs b/src/client/ernie.rs index ffe10f0..31ec536 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -18,52 +18,6 @@ use std::{env, sync::Mutex}; const API_BASE: &str = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1"; const ACCESS_TOKEN_URL: &str = "https://aip.baidubce.com/oauth/2.0/token"; -const MODELS: [(&str, &str, usize, isize); 7] = [ - // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/clntwmv7t - ( - "ernie-4.0-8k", - "/wenxinworkshop/chat/completions_pro", - 5120, - 2048, - ), - ( - "ernie-3.5-8k", - "/wenxinworkshop/chat/ernie-3.5-8k-0205", - 5120, - 2048, - ), - ( - "ernie-3.5-4k", - "/wenxinworkshop/chat/ernie-3.5-4k-0205", - 2048, - 2048, - ), - ( - "ernie-speed-8k", - "/wenxinworkshop/chat/ernie_speed", - 7168, - 2048, - ), - ( - "ernie-speed-128k", - "/wenxinworkshop/chat/ernie-speed-128k", - 124000, - 4096, - ), - ( - "ernie-lite-8k", - "/wenxinworkshop/chat/ernie-lite-8k", - 7168, - 2048, - ), - ( - "ernie-tiny-8k", - "/wenxinworkshop/chat/ernie-tiny-8k", - 7168, - 2048, - ), -]; - lazy_static! { static ref ACCESS_TOKEN: Mutex> = Mutex::new(None); } @@ -101,35 +55,38 @@ impl Client for ErnieClient { } impl ErnieClient { + list_models_fn!( + ErnieConfig, + [ + // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/clntwmv7t + ("ernie-4.0-8k", "text", 5120, 2048), + ("ernie-3.5-8k", "text", 5120, 2048), + ("ernie-3.5-4k", "text", 2048, 2048), + ("ernie-speed-8k", "text", 7168, 2048), + ("ernie-speed-128k", "text", 124000, 4096), + ("ernie-lite-8k", "text", 7168, 2048), + ("ernie-tiny-8k", "text", 7168, 2048), + ] + ); + pub const PROMPTS: [PromptType<'static>; 2] = [ ("api_key", "API Key:", true, PromptKind::String), ("secret_key", "Secret Key:", true, PromptKind::String), ]; - pub fn list_models(local_config: &ErnieConfig) -> Vec { - let client_name = Self::name(local_config); - if local_config.models.is_empty() { - MODELS - .into_iter() - .map(|(name, _, max_input_tokens, max_output_tokens)| { - Model::new(client_name, name) - .set_max_input_tokens(Some(max_input_tokens)) - .set_max_output_tokens(Some(max_output_tokens)) - }) // ERNIE tokenizer is different from cl100k_base - .collect() - } else { - Model::from_config(client_name, &local_config.models) - } - } - fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { let body = build_body(data, &self.model); - let model = &self.model.name; - let (_, chat_endpoint, _, _) = MODELS - .iter() - .find(|(v, _, _, _)| v == model) - .ok_or_else(|| anyhow!("Miss Model '{}'", self.model.id()))?; + let endpoint = match self.model.name.as_str() { + "ernie-4.0-8k" => "/wenxinworkshop/chat/completions_pro", + "ernie-3.5-8k" => "/wenxinworkshop/chat/ernie-3.5-8k-0205", + "ernie-3.5-4k" => "/wenxinworkshop/chat/ernie-3.5-4k-0205", + "ernie-speed-8k" => "/wenxinworkshop/chat/ernie_speed", + "ernie-speed-128k" => "/wenxinworkshop/chat/ernie-speed-128k", + "ernie-lite-8k" => "/wenxinworkshop/chat/ernie-lite-8k", + "ernie-tiny-8k" => "/wenxinworkshop/chat/ernie-tiny-8k", + _ => bail!("Miss Model '{}'", self.model.id()), + }; let access_token = ACCESS_TOKEN .lock() @@ -137,7 +94,7 @@ impl ErnieClient { .clone() .ok_or_else(|| anyhow!("Failed to load access token"))?; - let url = format!("{API_BASE}{chat_endpoint}?access_token={access_token}"); + let url = format!("{API_BASE}{endpoint}?access_token={access_token}"); debug!("Ernie Request: {url} {body}"); @@ -240,15 +197,14 @@ fn build_body(data: SendData, model: &Model) -> Value { "messages": messages, }); - if let Some(temperature) = temperature { - body["temperature"] = temperature.into(); + if let Some(v) = model.max_output_tokens { + body["max_output_tokens"] = v.into(); } - if let Some(top_p) = top_p { - body["top_p"] = top_p.into(); + if let Some(v) = temperature { + body["temperature"] = v.into(); } - - if let Some(max_output_tokens) = model.max_output_tokens { - body["max_output_tokens"] = max_output_tokens.into(); + if let Some(v) = top_p { + body["top_p"] = v.into(); } if stream { diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 22f32c5..b930d79 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -12,13 +12,6 @@ use serde::Deserialize; const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/"; -const MODELS: [(&str, usize, &str); 3] = [ - // https://ai.google.dev/models/gemini - ("gemini-1.0-pro-latest", 30720, "text"), - ("gemini-1.0-pro-vision-latest", 12288, "text,vision"), - ("gemini-1.5-pro-latest", 1048576, "text,vision"), -]; - #[derive(Debug, Clone, Deserialize, Default)] pub struct GeminiConfig { pub name: Option, @@ -50,7 +43,15 @@ impl Client for GeminiClient { } impl GeminiClient { - list_models_fn!(GeminiConfig, &MODELS); + list_models_fn!( + GeminiConfig, + [ + // https://ai.google.dev/models/gemini + ("gemini-1.0-pro-latest", "text", 30720), + ("gemini-1.0-pro-vision-latest", "text,vision", 12288), + ("gemini-1.5-pro-latest", "text,vision", 1048576), + ] + ); config_get_fn!(api_key, get_api_key); pub const PROMPTS: [PromptType<'static>; 1] = diff --git a/src/client/message.rs b/src/client/message.rs index f9978c5..9aff64a 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -114,6 +114,27 @@ pub struct ImageUrl { pub url: String, } +pub fn patch_system_message(messages: &mut Vec) { + if messages[0].role.is_system() { + let system_message = messages.remove(0); + if let (Some(message), MessageContent::Text(system_text)) = + (messages.get_mut(0), system_message.content) + { + if let MessageContent::Text(text) = message.content.clone() { + message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text)) + } + } + } +} + +pub fn extract_sytem_message(messages: &mut Vec) -> Option { + if messages[0].role.is_system() { + let system_message = messages.remove(0); + return Some(system_message.content.to_text()); + } + None +} + #[cfg(test)] mod tests { use super::*; diff --git a/src/client/mistral.rs b/src/client/mistral.rs index 9dbd831..23bf1f2 100644 --- a/src/client/mistral.rs +++ b/src/client/mistral.rs @@ -10,13 +10,6 @@ use serde::Deserialize; const API_URL: &str = "https://api.mistral.ai/v1/chat/completions"; -const MODELS: [(&str, usize, &str); 3] = [ - // https://docs.mistral.ai/platform/endpoints/ - ("open-mixtral-8x22b", 64000, "text"), - ("mistral-small-latest", 32000, "text"), - ("mistral-large-latest", 32000, "text"), -]; - #[derive(Debug, Clone, Deserialize)] pub struct MistralConfig { pub name: Option, @@ -29,7 +22,15 @@ pub struct MistralConfig { openai_compatible_client!(MistralClient); impl MistralClient { - list_models_fn!(MistralConfig, &MODELS); + list_models_fn!( + MistralConfig, + [ + // https://docs.mistral.ai/platform/endpoints/ + ("open-mixtral-8x22b", "text", 64000), + ("mistral-small-latest", "text", 32000), + ("mistral-large-latest", "text", 32000), + ] + ); config_get_fn!(api_key, get_api_key); pub const PROMPTS: [PromptType<'static>; 1] = diff --git a/src/client/model.rs b/src/client/model.rs index b97c244..3f2cbdd 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -49,17 +49,6 @@ impl Model { .collect() } - pub fn from_static(client_name: &str, models: &[(&str, usize, &str)]) -> Vec { - models - .iter() - .map(|(name, max_input_tokens, capabilities)| { - Model::new(client_name, name) - .set_capabilities((*capabilities).into()) - .set_max_input_tokens(Some(*max_input_tokens)) - }) - .collect() - } - pub fn find(models: &[Self], value: &str) -> Option { let mut model = None; let (client_name, model_name) = match value.split_once(':') { diff --git a/src/client/ollama.rs b/src/client/ollama.rs index 434e487..1394ced 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -179,15 +179,14 @@ fn build_body(data: SendData, model: &Model) -> Result { "options": {}, }); - if let Some(num_predict) = model.max_output_tokens { - body["options"]["num_predict"] = num_predict.into(); + if let Some(v) = model.max_output_tokens { + body["options"]["num_predict"] = v.into(); } - - if let Some(temperature) = temperature { - body["options"]["temperature"] = temperature.into(); + if let Some(v) = temperature { + body["options"]["temperature"] = v.into(); } - if let Some(top_p) = top_p { - body["options"]["top_p"] = top_p.into(); + if let Some(v) = top_p { + body["options"]["top_p"] = v.into(); } Ok(body) diff --git a/src/client/openai.rs b/src/client/openai.rs index 2c5d99b..37da878 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -12,18 +12,6 @@ use serde_json::{json, Value}; const API_BASE: &str = "https://api.openai.com/v1"; -const MODELS: [(&str, usize, &str); 8] = [ - // https://platform.openai.com/docs/models - ("gpt-3.5-turbo", 16385, "text"), - ("gpt-3.5-turbo-1106", 16385, "text"), - ("gpt-4-turbo", 128000, "text,vision"), - ("gpt-4-turbo-preview", 128000, "text"), - ("gpt-4-1106-preview", 128000, "text"), - ("gpt-4-vision-preview", 128000, "text,vision"), - ("gpt-4", 8192, "text"), - ("gpt-4-32k", 32768, "text"), -]; - #[derive(Debug, Clone, Deserialize, Default)] pub struct OpenAIConfig { pub name: Option, @@ -38,7 +26,20 @@ pub struct OpenAIConfig { openai_compatible_client!(OpenAIClient); impl OpenAIClient { - list_models_fn!(OpenAIConfig, &MODELS); + list_models_fn!( + OpenAIConfig, + [ + // https://platform.openai.com/docs/models + ("gpt-3.5-turbo", "text", 16385), + ("gpt-3.5-turbo-1106", "text", 16385), + ("gpt-4-turbo", "text,vision", 128000), + ("gpt-4-turbo-preview", "text", 128000), + ("gpt-4-1106-preview", "text", 128000), + ("gpt-4-vision-preview", "text,vision", 128000, 4096), + ("gpt-4", "text", 8192), + ("gpt-4-32k", "text", 32768), + ] + ); config_get_fn!(api_key, get_api_key); config_get_fn!(api_base, get_api_base); @@ -139,17 +140,14 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value { "messages": messages, }); - if let Some(max_tokens) = model.max_output_tokens { - body["max_tokens"] = max_tokens.into(); - } else if model.name == "gpt-4-vision-preview" { - // The default max_tokens of gpt-4-vision-preview is only 16, we need to make it larger - body["max_tokens"] = 4096.into(); + if let Some(v) = model.max_output_tokens { + body["max_tokens"] = v.into(); } - if let Some(temperature) = temperature { - body["temperature"] = temperature.into(); + if let Some(v) = temperature { + body["temperature"] = v.into(); } - if let Some(top_p) = top_p { - body["top_p"] = top_p.into(); + if let Some(v) = top_p { + body["top_p"] = v.into(); } if stream { body["stream"] = true.into(); diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index c5a72e0..52548d8 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -24,17 +24,6 @@ const API_URL: &str = const API_URL_VL: &str = "https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation"; -const MODELS: [(&str, usize, &str); 6] = [ - // https://help.aliyun.com/zh/dashscope/developer-reference/api-details - ("qwen-turbo", 6000, "text"), - ("qwen-plus", 30000, "text"), - ("qwen-max", 6000, "text"), - ("qwen-max-longcontext", 28000, "text"), - // https://help.aliyun.com/zh/dashscope/developer-reference/tongyi-qianwen-vl-plus-api - ("qwen-vl-plus", 0, "text,vision"), - ("qwen-vl-max", 0, "text,vision"), -]; - #[derive(Debug, Clone, Deserialize, Default)] pub struct QianwenConfig { pub name: Option, @@ -73,7 +62,19 @@ impl Client for QianwenClient { } impl QianwenClient { - list_models_fn!(QianwenConfig, &MODELS); + list_models_fn!( + QianwenConfig, + [ + // https://help.aliyun.com/zh/dashscope/developer-reference/api-details + ("qwen-turbo", "text", 6000), + ("qwen-plus", "text", 30000), + ("qwen-max", "text", 6000), + ("qwen-max-longcontext", "text", 28000), + // https://help.aliyun.com/zh/dashscope/developer-reference/tongyi-qianwen-vl-plus-api + ("qwen-vl-plus", "text,vision", 0), + ("qwen-vl-max", "text,vision", 0), + ] + ); config_get_fn!(api_key, get_api_key); pub const PROMPTS: [PromptType<'static>; 1] = @@ -211,15 +212,14 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool parameters["incremental_output"] = true.into(); } - if let Some(max_tokens) = model.max_output_tokens { - parameters["max_tokens"] = max_tokens.into(); + if let Some(v) = model.max_output_tokens { + parameters["max_tokens"] = v.into(); } - - if let Some(temperature) = temperature { - parameters["temperature"] = temperature.into(); + if let Some(v) = temperature { + parameters["temperature"] = v.into(); } - if let Some(top_p) = top_p { - parameters["top_p"] = top_p.into(); + if let Some(v) = top_p { + parameters["top_p"] = v.into(); } let body = json!({ diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index e0ae567..30aba75 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -13,13 +13,6 @@ use serde::Deserialize; use serde_json::{json, Value}; use std::path::PathBuf; -const MODELS: [(&str, usize, &str); 3] = [ - // https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models - ("gemini-1.0-pro", 24568, "text"), - ("gemini-1.0-pro-vision", 14336, "text,vision"), - ("gemini-1.5-pro-preview-0409", 1000000, "text,vision"), -]; - static mut ACCESS_TOKEN: (String, i64) = (String::new(), 0); // safe under linear operation #[derive(Debug, Clone, Deserialize, Default)] @@ -56,7 +49,15 @@ impl Client for VertexAIClient { } impl VertexAIClient { - list_models_fn!(VertexAIConfig, &MODELS); + list_models_fn!( + VertexAIConfig, + [ + // https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models + ("gemini-1.0-pro", "text", 24568), + ("gemini-1.0-pro-vision", "text,vision", 14336), + ("gemini-1.5-pro-preview-0409", "text,vision", 1000000), + ] + ); config_get_fn!(api_base, get_api_base); pub const PROMPTS: [PromptType<'static>; 1] = @@ -216,16 +217,14 @@ pub(crate) fn build_body( ]); } - if let Some(max_output_tokens) = model.max_output_tokens { - body["generationConfig"]["maxOutputTokens"] = max_output_tokens.into(); + if let Some(v) = model.max_output_tokens { + body["generationConfig"]["maxOutputTokens"] = v.into(); } - - if let Some(temperature) = temperature { - body["generationConfig"]["temperature"] = temperature.into(); + if let Some(v) = temperature { + body["generationConfig"]["temperature"] = v.into(); } - - if let Some(top_p) = top_p { - body["generationConfig"]["topP"] = top_p.into(); + if let Some(v) = top_p { + body["generationConfig"]["topP"] = v.into(); } Ok(body) -- cgit v1.2.3