From d1aafa11153ab689c21c2c57c47da52337d8e8d1 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 23 Apr 2024 16:46:48 +0800 Subject: feat: customize model's max_output_tokens (#428) --- src/client/cohere.rs | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) (limited to 'src/client/cohere.rs') diff --git a/src/client/cohere.rs b/src/client/cohere.rs index a92e238..bfea105 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -67,7 +67,7 @@ impl CohereClient { fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { let api_key = self.get_api_key().ok(); - let body = build_body(data, self.model.name.clone())?; + let body = build_body(data, &self.model)?; let url = API_URL; @@ -131,7 +131,7 @@ fn check_error(data: &Value) -> Result<()> { } } -pub(crate) fn build_body(data: SendData, model: String) -> Result { +pub(crate) fn build_body(data: SendData, model: &Model) -> Result { let SendData { mut messages, temperature, @@ -179,9 +179,13 @@ pub(crate) fn build_body(data: SendData, model: String) -> Result { let message = message["message"].as_str().unwrap_or_default(); let mut body = json!({ - "model": model, + "model": &model.name, "message": message, }); + + if let Some(max_tokens) = model.max_output_tokens { + body["max_tokens"] = max_tokens.into(); + } if !messages.is_empty() { body["chat_history"] = messages.into(); -- cgit v1.2.3