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/client/replicate.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/client/replicate.rs')
| -rw-r--r-- | src/client/replicate.rs | 26 |
1 files changed, 15 insertions, 11 deletions
diff --git a/src/client/replicate.rs b/src/client/replicate.rs index 34cfd94..c0a77c2 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -1,9 +1,9 @@ use std::time::Duration; use super::{ - catch_error, generate_prompt, smart_prompt_format, sse_stream, Client, CompletionDetails, - ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, ReplicateClient, SendData, - SsMmessage, SseHandler, + catch_error, generate_prompt, smart_prompt_format, sse_stream, Client, CompletionOutput, + ExtraConfig, Model, ModelData, PromptAction, PromptKind, ReplicateClient, SendData, SsMmessage, + SseHandler, }; use anyhow::{anyhow, Result}; @@ -19,7 +19,7 @@ pub struct ReplicateConfig { pub name: Option<String>, pub api_key: Option<String>, #[serde(default)] - pub models: Vec<ModelConfig>, + pub models: Vec<ModelData>, pub extra: Option<ExtraConfig>, } @@ -37,7 +37,7 @@ impl ReplicateClient { ) -> Result<RequestBuilder> { let body = build_body(data, &self.model)?; - let url = format!("{API_BASE}/models/{}/predictions", self.model.name); + let url = format!("{API_BASE}/models/{}/predictions", self.model.name()); debug!("Replicate Request: {url} {body}"); @@ -55,7 +55,7 @@ impl Client for ReplicateClient { &self, client: &ReqwestClient, data: SendData, - ) -> Result<(String, CompletionDetails)> { + ) -> Result<CompletionOutput> { let api_key = self.get_api_key()?; let builder = self.request_builder(client, data, &api_key)?; send_message(client, builder, &api_key).await @@ -77,7 +77,7 @@ async fn send_message( client: &ReqwestClient, builder: RequestBuilder, api_key: &str, -) -> Result<(String, CompletionDetails)> { +) -> Result<CompletionOutput> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -96,6 +96,7 @@ async fn send_message( .await? .json() .await?; + debug!("non-stream-data: {prediction_data}"); let err = || anyhow!("Invalid response data: {prediction_data}"); let status = prediction_data["status"].as_str().ok_or_else(err)?; if status == "succeeded" { @@ -138,10 +139,11 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> { messages, temperature, top_p, + functions: _, stream, } = data; - let prompt = generate_prompt(&messages, smart_prompt_format(&model.name))?; + let prompt = generate_prompt(&messages, smart_prompt_format(model.name()))?; let mut input = json!({ "prompt": prompt, @@ -170,7 +172,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> { Ok(body) } -fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> { +fn extract_completion(data: &Value) -> Result<CompletionOutput> { let text = data["output"] .as_array() .map(|parts| { @@ -182,11 +184,13 @@ fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> { }) .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - let details = CompletionDetails { + let output = CompletionOutput { + text: text.to_string(), + tool_calls: vec![], id: data["id"].as_str().map(|v| v.to_string()), input_tokens: data["metrics"]["input_token_count"].as_u64(), output_tokens: data["metrics"]["output_token_count"].as_u64(), }; - Ok((text.to_string(), details)) + Ok(output) } |
