diff options
Diffstat (limited to 'src/client/replicate.rs')
| -rw-r--r-- | src/client/replicate.rs | 24 |
1 files changed, 12 insertions, 12 deletions
diff --git a/src/client/replicate.rs b/src/client/replicate.rs index 53097e6..d6ca401 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -16,7 +16,7 @@ pub struct ReplicateConfig { pub api_key: Option<String>, #[serde(default)] pub models: Vec<ModelData>, - pub patch: Option<ModelPatch>, + pub patch: Option<RequestPatch>, pub extra: Option<ExtraConfig>, } @@ -26,22 +26,20 @@ impl ReplicateClient { pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; - fn chat_completions_builder( + fn prepare_chat_completions( &self, - client: &ReqwestClient, data: ChatCompletionsData, api_key: &str, - ) -> Result<RequestBuilder> { - let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_chat_completions_body(&mut body); - + ) -> Result<RequestData> { let url = format!("{API_BASE}/models/{}/predictions", self.model.name()); - debug!("Replicate Request: {url} {body}"); + let body = build_chat_completions_body(data, &self.model)?; + + let mut request_data = RequestData::new(url, body); - let builder = client.post(url).bearer_auth(api_key).json(&body); + request_data.bearer_auth(api_key); - Ok(builder) + Ok(request_data) } } @@ -55,7 +53,8 @@ impl Client for ReplicateClient { data: ChatCompletionsData, ) -> Result<ChatCompletionsOutput> { let api_key = self.get_api_key()?; - let builder = self.chat_completions_builder(client, data, &api_key)?; + let request_data = self.prepare_chat_completions(data, &api_key)?; + let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); chat_completions(client, builder, &api_key).await } @@ -66,7 +65,8 @@ impl Client for ReplicateClient { data: ChatCompletionsData, ) -> Result<()> { let api_key = self.get_api_key()?; - let builder = self.chat_completions_builder(client, data, &api_key)?; + let request_data = self.prepare_chat_completions(data, &api_key)?; + let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); chat_completions_streaming(client, builder, handler).await } } |
