From ba3bcfd67c1d6fea5d3d3c5908c975682ee7909b Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 22 May 2024 21:29:23 +0800 Subject: feat: allow patching req body with client config (#534) --- src/client/replicate.rs | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) (limited to 'src/client/replicate.rs') diff --git a/src/client/replicate.rs b/src/client/replicate.rs index c0a77c2..c549ae8 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -1,9 +1,7 @@ -use std::time::Duration; - use super::{ - catch_error, generate_prompt, smart_prompt_format, sse_stream, Client, CompletionOutput, - ExtraConfig, Model, ModelData, PromptAction, PromptKind, ReplicateClient, SendData, SsMmessage, - SseHandler, + catch_error, prompt_format::*, sse_stream, Client, CompletionOutput, ExtraConfig, + Model, ModelData, ModelPatches, PromptAction, PromptKind, ReplicateClient, SendData, + SsMmessage, SseHandler, }; use anyhow::{anyhow, Result}; @@ -11,6 +9,7 @@ use async_trait::async_trait; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; +use std::time::Duration; const API_BASE: &str = "https://api.replicate.com/v1"; @@ -20,6 +19,7 @@ pub struct ReplicateConfig { pub api_key: Option, #[serde(default)] pub models: Vec, + pub patches: Option, pub extra: Option, } @@ -35,7 +35,8 @@ impl ReplicateClient { data: SendData, api_key: &str, ) -> Result { - let body = build_body(data, &self.model)?; + let mut body = build_body(data, &self.model)?; + self.patch_request_body(&mut body); let url = format!("{API_BASE}/models/{}/predictions", self.model.name()); -- cgit v1.2.3