diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-18 19:06:21 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-18 19:06:21 +0800 |
| commit | b4a40e3fedb438570770a224b890ea24f6e660a9 (patch) | |
| tree | 344b96102da7cbedf1034d023aa82599940388b1 /src/serve.rs | |
| parent | 1348a62e5f8bc140a7218fbfe1b73f990ab16101 (diff) | |
| download | aichat-b4a40e3fedb438570770a224b890ea24f6e660a9.tar.gz | |
feat: support function calling (#514)
* feat: support function calling
* fix on Windows OS
* implement multi-steps function calling
* fix on Windows OS
* add error for client not support function calling
* refactor message data structure and make claude client supporting function calling
* support reuse previous call results
* improve error handling for function calling
* use prefix `may_` as indicator for `execute` type fucntions
Diffstat (limited to 'src/serve.rs')
| -rw-r--r-- | src/serve.rs | 51 |
1 files changed, 28 insertions, 23 deletions
diff --git a/src/serve.rs b/src/serve.rs index a4f7d02..0e7bda3 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -1,7 +1,7 @@ use crate::{ client::{ - init_client, list_models, ClientConfig, CompletionDetails, Message, Model, SendData, - SseEvent, SseHandler, + init_client, list_models, ClientConfig, CompletionOutput, Message, Model, ModelData, + SendData, SseEvent, SseHandler, }, config::{Config, GlobalConfig, Role}, utils::create_abort_signal, @@ -78,7 +78,7 @@ impl Server { let roles = config.roles.clone(); let mut models = list_models(&config); let mut default_model = model.clone(); - default_model.name = DEFAULT_MODEL_NAME.into(); + default_model.data_mut().name = DEFAULT_MODEL_NAME.into(); models.insert(0, &default_model); let models: Vec<Value> = models .into_iter() @@ -89,14 +89,25 @@ impl Server { } else { model.id() }; + let ModelData { + max_input_tokens, + max_output_tokens, + pass_max_tokens, + input_price, + output_price, + supports_vision, + supports_function_calling, + .. + } = model.data(); json!({ "id": id, - "max_input_tokens": model.max_input_tokens, - "max_output_tokens": model.max_output_tokens, - "pass_max_tokens": model.pass_max_tokens, - "input_price": model.input_price, - "output_price": model.output_price, - "supports_vision": model.supports_vision(), + "max_input_tokens": max_input_tokens, + "max_output_tokens": max_output_tokens, + "pass_max_tokens": pass_max_tokens, + "input_price": input_price, + "output_price": output_price, + "supports_vision": supports_vision, + "supports_function_calling": supports_function_calling, }) }) .collect(); @@ -263,6 +274,7 @@ impl Server { messages, temperature, top_p, + functions: None, stream, }; @@ -338,7 +350,7 @@ impl Server { .body(BodyExt::boxed(StreamBody::new(stream)))?; Ok(res) } else { - let (content, details) = client.send_message_inner(&http_client, send_data).await?; + let output = client.send_message_inner(&http_client, send_data).await?; let res = Response::builder() .header("Content-Type", "application/json") .body( @@ -346,8 +358,7 @@ impl Server { &completion_id, &model_name, created, - &content, - &details, + &output, )) .boxed(), )?; @@ -439,16 +450,10 @@ fn create_frame(id: &str, model: &str, created: i64, content: &str, done: bool) Frame::data(Bytes::from(output)) } -fn ret_non_stream( - id: &str, - model: &str, - created: i64, - content: &str, - details: &CompletionDetails, -) -> Bytes { - let id = details.id.as_deref().unwrap_or(id); - let input_tokens = details.input_tokens.unwrap_or_default(); - let output_tokens = details.output_tokens.unwrap_or_default(); +fn ret_non_stream(id: &str, model: &str, created: i64, output: &CompletionOutput) -> Bytes { + let id = output.id.as_deref().unwrap_or(id); + let input_tokens = output.input_tokens.unwrap_or_default(); + let output_tokens = output.output_tokens.unwrap_or_default(); let total_tokens = input_tokens + output_tokens; let res_body = json!({ "id": id, @@ -460,7 +465,7 @@ fn ret_non_stream( "index": 0, "message": { "role": "assistant", - "content": content, + "content": output.text, }, "logprobs": null, "finish_reason": "stop", |
