kimun_notes/server_client/
mod.rs1use std::collections::HashMap;
16use std::time::Duration;
17
18pub mod dto;
19pub mod observer;
20pub mod reconcile;
21
22pub mod sync;
23
24use async_trait::async_trait;
25use dto::{
26 AnswerResult, DeleteRequest, EmbeddingsResponse, Health, HistoryTurn, IndexDocsRequest,
27 JobAccepted, JobStatus, QueryRequest, WireDoc,
28};
29
30pub use dto::{ChunkResult, WireSection};
31pub use observer::{DirtyOp, DirtySet, RagObserver};
32pub use reconcile::{ReconcilePlan, diff as reconcile_diff};
33
34#[derive(Debug, thiserror::Error)]
35pub enum RagError {
36 #[error("http error: {0}")]
37 Http(#[from] reqwest::Error),
38 #[error("server returned {status}: {body}")]
39 Status { status: u16, body: String },
40 #[error("{0}")]
43 Protocol(String),
44}
45
46impl RagError {
47 pub fn is_auth(&self) -> bool {
51 matches!(
52 self,
53 RagError::Status {
54 status: 401 | 403,
55 ..
56 }
57 )
58 }
59}
60
61#[allow(clippy::double_must_use)]
65#[async_trait]
66pub trait RagTransport: Send + Sync {
67 async fn push_docs(&self, docs: Vec<WireDoc>) -> Result<(), RagError>;
68 async fn delete_paths(&self, paths: Vec<String>) -> Result<(), RagError>;
69 async fn server_hashes(&self) -> Result<HashMap<String, String>, RagError>;
70}
71
72#[async_trait]
73impl RagTransport for RagClient {
74 async fn push_docs(&self, docs: Vec<WireDoc>) -> Result<(), RagError> {
75 RagClient::push_docs(self, docs).await.map(|_job_id| ())
78 }
79 async fn delete_paths(&self, paths: Vec<String>) -> Result<(), RagError> {
80 RagClient::delete_paths(self, paths).await
81 }
82 async fn server_hashes(&self) -> Result<HashMap<String, String>, RagError> {
83 RagClient::server_hashes(self).await
84 }
85}
86
87pub fn hash_string(hash: u64) -> String {
92 hash.to_string()
93}
94
95#[derive(Debug, Clone, Copy)]
98pub enum ContextSize {
99 Small,
100 Medium,
101 Large,
102}
103
104impl ContextSize {
105 fn as_str(self) -> &'static str {
106 match self {
107 ContextSize::Small => "small",
108 ContextSize::Medium => "medium",
109 ContextSize::Large => "large",
110 }
111 }
112}
113
114const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
118
119const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
122
123const PUSH_TIMEOUT: Duration = Duration::from_secs(120);
127
128fn shared_http() -> reqwest::Client {
132 static HTTP: std::sync::OnceLock<reqwest::Client> = std::sync::OnceLock::new();
133 HTTP.get_or_init(|| {
134 reqwest::Client::builder()
135 .connect_timeout(CONNECT_TIMEOUT)
136 .timeout(REQUEST_TIMEOUT)
137 .build()
138 .expect("build HTTP client")
141 })
142 .clone()
143}
144
145#[derive(Clone)]
147pub struct RagClient {
148 http: reqwest::Client,
149 base_url: String,
150 token: Option<String>,
151 vault_id: String,
152}
153
154impl RagClient {
155 pub fn new(
158 base_url: impl Into<String>,
159 token: Option<String>,
160 vault_id: impl Into<String>,
161 ) -> Self {
162 let base_url = base_url.into().trim_end_matches('/').to_string();
163 Self {
164 http: shared_http(),
165 base_url,
166 token,
167 vault_id: vault_id.into(),
168 }
169 }
170
171 fn url(&self, path: &str) -> String {
172 format!("{}{}", self.base_url, path)
173 }
174
175 fn auth(&self, req: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
177 match &self.token {
178 Some(token) => req.bearer_auth(token),
179 None => req,
180 }
181 }
182
183 async fn ok(resp: reqwest::Response) -> Result<reqwest::Response, RagError> {
186 if resp.status().is_success() {
187 Ok(resp)
188 } else {
189 let status = resp.status().as_u16();
190 let body = resp.text().await.unwrap_or_default();
191 Err(RagError::Status { status, body })
192 }
193 }
194
195 pub async fn health(&self) -> Result<Health, RagError> {
197 let resp = self.auth(self.http.get(self.url("/health"))).send().await?;
198 Ok(Self::ok(resp).await?.json::<Health>().await?)
199 }
200
201 pub async fn push_docs(&self, docs: Vec<WireDoc>) -> Result<String, RagError> {
203 let body = IndexDocsRequest {
204 vault_id: self.vault_id.clone(),
205 docs,
206 };
207 let resp = self
208 .auth(self.http.post(self.url("/api/index/docs")).json(&body))
209 .timeout(PUSH_TIMEOUT)
210 .send()
211 .await?;
212 Ok(Self::ok(resp).await?.json::<JobAccepted>().await?.job_id)
213 }
214
215 pub async fn delete_paths(&self, paths: Vec<String>) -> Result<(), RagError> {
217 let body = DeleteRequest {
218 vault_id: self.vault_id.clone(),
219 paths,
220 };
221 let resp = self
222 .auth(self.http.post(self.url("/api/index/delete")).json(&body))
223 .send()
224 .await?;
225 Self::ok(resp).await?;
226 Ok(())
227 }
228
229 pub async fn server_hashes(&self) -> Result<HashMap<String, String>, RagError> {
235 let path = format!("/api/collections/{}/hashes", self.vault_id);
236 let resp = self.auth(self.http.get(self.url(&path))).send().await?;
237 Ok(Self::ok(resp)
238 .await?
239 .json::<HashMap<String, String>>()
240 .await?)
241 }
242
243 pub async fn search(
246 &self,
247 query: &str,
248 context_size: Option<ContextSize>,
249 ) -> Result<Vec<ChunkResult>, RagError> {
250 let body = QueryRequest {
251 vault_id: self.vault_id.clone(),
252 query: query.to_string(),
253 context_size: context_size.map(|c| c.as_str().to_string()),
254 history: vec![],
255 };
256 let resp = self
257 .auth(self.http.post(self.url("/api/embeddings")).json(&body))
258 .send()
259 .await?;
260 Ok(Self::ok(resp)
261 .await?
262 .json::<EmbeddingsResponse>()
263 .await?
264 .chunks)
265 }
266
267 pub async fn ask(
270 &self,
271 query: &str,
272 history: &[(String, String)],
273 context_size: Option<ContextSize>,
274 ) -> Result<AnswerResult, RagError> {
275 let body = QueryRequest {
276 vault_id: self.vault_id.clone(),
277 query: query.to_string(),
278 context_size: context_size.map(|c| c.as_str().to_string()),
279 history: history
280 .iter()
281 .map(|(q, a)| HistoryTurn {
282 question: q.clone(),
283 answer: a.clone(),
284 })
285 .collect(),
286 };
287 let resp = self
288 .auth(self.http.post(self.url("/api/answer")).json(&body))
289 .send()
290 .await?;
291 let job_id = Self::ok(resp).await?.json::<JobAccepted>().await?.job_id;
292 self.poll_answer(&job_id).await
293 }
294
295 const ANSWER_POLL_ATTEMPTS: u32 = 720;
302
303 async fn poll_answer(&self, job_id: &str) -> Result<AnswerResult, RagError> {
305 let path = format!("/api/job/{job_id}");
306 let mut consecutive_errors = 0u32;
307 for _ in 0..Self::ANSWER_POLL_ATTEMPTS {
308 let status = match self.poll_once(&path).await {
312 Ok(s) => {
313 consecutive_errors = 0;
314 s
315 }
316 Err(e) => {
317 consecutive_errors += 1;
318 if consecutive_errors >= 15 {
319 return Err(e);
320 }
321 tokio::time::sleep(std::time::Duration::from_secs(1)).await;
322 continue;
323 }
324 };
325 match status.status.as_str() {
326 "completed" => {
327 let result = status
328 .result
329 .ok_or_else(|| RagError::Protocol("completed job had no result".into()))?;
330 return serde_json::from_value::<AnswerResult>(result)
331 .map_err(|e| RagError::Protocol(format!("bad answer result: {e}")));
332 }
333 "failed" => {
334 return Err(RagError::Protocol(
335 status.error.unwrap_or_else(|| "answer job failed".into()),
336 ));
337 }
338 _ => tokio::time::sleep(std::time::Duration::from_secs(1)).await,
339 }
340 }
341 Err(RagError::Protocol("answer job timed out".into()))
342 }
343
344 async fn poll_once(&self, path: &str) -> Result<JobStatus, RagError> {
346 let resp = self.auth(self.http.get(self.url(path))).send().await?;
347 Ok(Self::ok(resp).await?.json::<JobStatus>().await?)
348 }
349}
350
351#[cfg(test)]
352mod tests {
353 use super::*;
354
355 #[test]
356 fn is_auth_matches_only_credential_rejections() {
357 let status = |status| RagError::Status {
358 status,
359 body: String::new(),
360 };
361 assert!(status(401).is_auth());
362 assert!(status(403).is_auth());
363 assert!(!status(500).is_auth());
364 assert!(!status(404).is_auth());
365 assert!(!RagError::Protocol("boom".into()).is_auth());
366 }
367}