summaryrefslogtreecommitdiffstats
path: root/src/client/ollama.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-28 06:04:36 +0800
committerGitHub <noreply@github.com>2024-07-28 06:04:36 +0800
commit0e740d81e94505bd57036755abaaecb12c3b26e3 (patch)
tree49000370fb12e4e5e1f1bd4d145104f2f502aa16 /src/client/ollama.rs
parentf5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff)
downloadaichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/ollama.rs')
-rw-r--r--src/client/ollama.rs75
1 files changed, 43 insertions, 32 deletions
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 5c26b99..4c2f344 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -31,53 +31,63 @@ impl OllamaClient {
PromptKind::Integer,
),
];
+}
- fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> {
- let api_base = self.get_api_base()?;
- let api_auth = self.get_api_auth().ok();
+impl_client_trait!(
+ OllamaClient,
+ (
+ prepare_chat_completions,
+ chat_completions,
+ chat_completions_streaming
+ ),
+ (prepare_embeddings, embeddings),
+ (noop_prepare_rerank, noop_rerank),
+);
- let url = format!("{api_base}/api/chat");
+fn prepare_chat_completions(
+ self_: &OllamaClient,
+ data: ChatCompletionsData,
+) -> Result<RequestData> {
+ let api_base = self_.get_api_base()?;
+ let api_auth = self_.get_api_auth().ok();
- let body = build_chat_completions_body(data, &self.model)?;
+ let url = format!("{api_base}/api/chat");
- let mut request_data = RequestData::new(url, body);
+ let body = build_chat_completions_body(data, &self_.model)?;
- if let Some(api_auth) = api_auth {
- request_data.header("Authorization", api_auth)
- }
+ let mut request_data = RequestData::new(url, body);
- Ok(request_data)
+ if let Some(api_auth) = api_auth {
+ request_data.header("Authorization", api_auth)
}
- fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> {
- let api_base = self.get_api_base()?;
- let api_auth = self.get_api_auth().ok();
+ Ok(request_data)
+}
- let url = format!("{api_base}/api/embed");
+fn prepare_embeddings(self_: &OllamaClient, data: EmbeddingsData) -> Result<RequestData> {
+ let api_base = self_.get_api_base()?;
+ let api_auth = self_.get_api_auth().ok();
- let body = json!({
- "model": self.model.name(),
- "input": data.texts,
- });
+ let url = format!("{api_base}/api/embed");
- let mut request_data = RequestData::new(url, body);
+ let body = json!({
+ "model": self_.model.name(),
+ "input": data.texts,
+ });
- if let Some(api_auth) = api_auth {
- request_data.header("Authorization", api_auth)
- }
+ let mut request_data = RequestData::new(url, body);
- Ok(request_data)
+ if let Some(api_auth) = api_auth {
+ request_data.header("Authorization", api_auth)
}
-}
-impl_client_trait!(
- OllamaClient,
- chat_completions,
- chat_completions_streaming,
- embeddings
-);
+ Ok(request_data)
+}
-async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
+async fn chat_completions(
+ builder: RequestBuilder,
+ _model: &Model,
+) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
let status = res.status();
let data = res.json().await?;
@@ -92,6 +102,7 @@ async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutp
async fn chat_completions_streaming(
builder: RequestBuilder,
handler: &mut SseHandler,
+ _model: &Model,
) -> Result<()> {
let res = builder.send().await?;
let status = res.status();
@@ -120,7 +131,7 @@ async fn chat_completions_streaming(
Ok(())
}
-async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
+async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result<EmbeddingsOutput> {
let res = builder.send().await?;
let status = res.status();
let data = res.json().await?;