summaryrefslogtreecommitdiffstats
path: root/src/client/replicate.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/replicate.rs')
-rw-r--r--src/client/replicate.rs24
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
}
}