1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize};
4use serde_json::{to_vec, Value};
5
6use crate::{
7 client::Client,
8 errors::Error,
9 request::{HttpClient, Method},
10};
11
12#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
14#[serde(rename_all = "camelCase")]
15pub struct ChatWorkspace {
16 pub uid: String,
17}
18
19#[derive(Debug, Clone, Deserialize, Serialize)]
21#[serde(rename_all = "camelCase")]
22pub struct ChatWorkspacesResults {
23 pub results: Vec<ChatWorkspace>,
24 pub offset: u32,
25 pub limit: u32,
26 pub total: u32,
27}
28
29#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
31#[serde(rename_all = "camelCase")]
32pub struct ChatPrompts {
33 #[serde(skip_serializing_if = "Option::is_none")]
34 pub system: Option<String>,
35 #[serde(skip_serializing_if = "Option::is_none")]
36 pub search_description: Option<String>,
37 #[serde(rename = "searchQParam", skip_serializing_if = "Option::is_none")]
38 pub search_q_param: Option<String>,
39 #[serde(
40 rename = "searchIndexUidParam",
41 skip_serializing_if = "Option::is_none"
42 )]
43 pub search_index_uid_param: Option<String>,
44 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
46 pub extra: BTreeMap<String, String>,
47}
48
49impl ChatPrompts {
50 #[must_use]
51 pub fn new() -> Self {
52 Self::default()
53 }
54
55 pub fn set_system(&mut self, value: impl Into<String>) -> &mut Self {
56 self.system = Some(value.into());
57 self
58 }
59
60 pub fn set_search_description(&mut self, value: impl Into<String>) -> &mut Self {
61 self.search_description = Some(value.into());
62 self
63 }
64
65 pub fn set_search_q_param(&mut self, value: impl Into<String>) -> &mut Self {
66 self.search_q_param = Some(value.into());
67 self
68 }
69
70 pub fn set_search_index_uid_param(&mut self, value: impl Into<String>) -> &mut Self {
71 self.search_index_uid_param = Some(value.into());
72 self
73 }
74
75 pub fn insert(&mut self, key: impl Into<String>, value: impl Into<String>) -> &mut Self {
76 self.extra.insert(key.into(), value.into());
77 self
78 }
79}
80
81#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
83#[serde(rename_all = "camelCase")]
84pub struct ChatWorkspaceSettings {
85 #[serde(skip_serializing_if = "Option::is_none")]
86 pub source: Option<String>,
87 #[serde(skip_serializing_if = "Option::is_none")]
88 pub org_id: Option<String>,
89 #[serde(skip_serializing_if = "Option::is_none")]
90 pub project_id: Option<String>,
91 #[serde(skip_serializing_if = "Option::is_none")]
92 pub api_version: Option<String>,
93 #[serde(skip_serializing_if = "Option::is_none")]
94 pub deployment_id: Option<String>,
95 #[serde(skip_serializing_if = "Option::is_none")]
96 pub base_url: Option<String>,
97 #[serde(skip_serializing_if = "Option::is_none")]
98 pub api_key: Option<String>,
99 #[serde(skip_serializing_if = "Option::is_none")]
100 pub prompts: Option<ChatPrompts>,
101}
102
103impl ChatWorkspaceSettings {
104 #[must_use]
105 pub fn new() -> Self {
106 Self::default()
107 }
108
109 pub fn set_source(&mut self, source: impl Into<String>) -> &mut Self {
110 self.source = Some(source.into());
111 self
112 }
113
114 pub fn set_org_id(&mut self, org_id: impl Into<String>) -> &mut Self {
115 self.org_id = Some(org_id.into());
116 self
117 }
118
119 pub fn set_project_id(&mut self, project_id: impl Into<String>) -> &mut Self {
120 self.project_id = Some(project_id.into());
121 self
122 }
123
124 pub fn set_api_version(&mut self, api_version: impl Into<String>) -> &mut Self {
125 self.api_version = Some(api_version.into());
126 self
127 }
128
129 pub fn set_deployment_id(&mut self, deployment_id: impl Into<String>) -> &mut Self {
130 self.deployment_id = Some(deployment_id.into());
131 self
132 }
133
134 pub fn set_base_url(&mut self, base_url: impl Into<String>) -> &mut Self {
135 self.base_url = Some(base_url.into());
136 self
137 }
138
139 pub fn set_api_key(&mut self, api_key: impl Into<String>) -> &mut Self {
140 self.api_key = Some(api_key.into());
141 self
142 }
143
144 pub fn set_prompts(&mut self, prompts: impl Into<ChatPrompts>) -> &mut Self {
145 self.prompts = Some(prompts.into());
146 self
147 }
148}
149
150#[derive(Debug, Serialize)]
152pub struct ChatWorkspacesQuery<'a, Http: HttpClient> {
153 #[serde(skip_serializing)]
154 pub client: &'a Client<Http>,
155 #[serde(skip_serializing_if = "Option::is_none")]
156 pub offset: Option<usize>,
157 #[serde(skip_serializing_if = "Option::is_none")]
158 pub limit: Option<usize>,
159}
160
161impl<'a, Http: HttpClient> ChatWorkspacesQuery<'a, Http> {
162 #[must_use]
163 pub fn new(client: &'a Client<Http>) -> Self {
164 Self {
165 client,
166 offset: None,
167 limit: None,
168 }
169 }
170
171 pub fn with_offset(&mut self, offset: usize) -> &mut Self {
172 self.offset = Some(offset);
173 self
174 }
175
176 pub fn with_limit(&mut self, limit: usize) -> &mut Self {
177 self.limit = Some(limit);
178 self
179 }
180
181 pub async fn execute(&self) -> Result<ChatWorkspacesResults, Error> {
182 self.client.list_chat_workspaces_with(self).await
183 }
184}
185
186impl<Http: HttpClient> Client<Http> {
187 pub async fn list_chat_workspaces(&self) -> Result<ChatWorkspacesResults, Error> {
189 self.http_client
190 .request::<(), (), ChatWorkspacesResults>(
191 &format!("{}/chats", self.host),
192 Method::Get { query: () },
193 200,
194 )
195 .await
196 }
197
198 pub async fn list_chat_workspaces_with(
200 &self,
201 query: &ChatWorkspacesQuery<'_, Http>,
202 ) -> Result<ChatWorkspacesResults, Error> {
203 self.http_client
204 .request::<&ChatWorkspacesQuery<'_, Http>, (), ChatWorkspacesResults>(
205 &format!("{}/chats", self.host),
206 Method::Get { query },
207 200,
208 )
209 .await
210 }
211
212 pub async fn get_chat_workspace(&self, uid: impl AsRef<str>) -> Result<ChatWorkspace, Error> {
214 self.http_client
215 .request::<(), (), ChatWorkspace>(
216 &format!("{}/chats/{}", self.host, uid.as_ref()),
217 Method::Get { query: () },
218 200,
219 )
220 .await
221 }
222
223 pub async fn get_chat_workspace_settings(
225 &self,
226 uid: impl AsRef<str>,
227 ) -> Result<ChatWorkspaceSettings, Error> {
228 self.http_client
229 .request::<(), (), ChatWorkspaceSettings>(
230 &format!("{}/chats/{}/settings", self.host, uid.as_ref()),
231 Method::Get { query: () },
232 200,
233 )
234 .await
235 }
236
237 pub async fn update_chat_workspace_settings(
239 &self,
240 uid: impl AsRef<str>,
241 settings: &ChatWorkspaceSettings,
242 ) -> Result<ChatWorkspaceSettings, Error> {
243 self.http_client
244 .request::<(), &ChatWorkspaceSettings, ChatWorkspaceSettings>(
245 &format!("{}/chats/{}/settings", self.host, uid.as_ref()),
246 Method::Patch {
247 query: (),
248 body: settings,
249 },
250 200,
251 )
252 .await
253 }
254
255 pub async fn reset_chat_workspace_settings(
257 &self,
258 uid: impl AsRef<str>,
259 ) -> Result<ChatWorkspaceSettings, Error> {
260 self.http_client
261 .request::<(), (), ChatWorkspaceSettings>(
262 &format!("{}/chats/{}/settings", self.host, uid.as_ref()),
263 Method::Delete { query: () },
264 200,
265 )
266 .await
267 }
268}
269
270#[cfg(feature = "reqwest")]
271impl Client<crate::reqwest::ReqwestClient> {
272 pub async fn stream_chat_completion<S: Serialize + ?Sized>(
274 &self,
275 uid: impl AsRef<str>,
276 body: &S,
277 ) -> Result<reqwest::Response, Error> {
278 let request = self.build_stream_chat_request(uid.as_ref(), body)?;
279
280 let response = self.http_client.inner().execute(request).await?;
281
282 let status = response.status();
283 if !status.is_success() {
284 let url = response.url().to_string();
285 let mut body = response.text().await?;
286 if body.is_empty() {
287 body = "null".to_string();
288 }
289 let err =
290 match crate::request::parse_response::<Value>(status.as_u16(), 200, &body, url) {
291 Ok(_) => unreachable!("parse_response succeeded on a non-successful status"),
292 Err(err) => err,
293 };
294 return Err(err);
295 }
296
297 Ok(response)
298 }
299
300 fn build_stream_chat_request<S: Serialize + ?Sized>(
301 &self,
302 uid: &str,
303 body: &S,
304 ) -> Result<reqwest::Request, Error> {
305 use reqwest::header::{HeaderValue, ACCEPT, AUTHORIZATION, CONTENT_TYPE};
306
307 let payload = to_vec(body).map_err(Error::ParseError)?;
308
309 let mut request = self
310 .http_client
311 .inner()
312 .post(format!("{}/chats/{}/chat/completions", self.host, uid))
313 .header(ACCEPT, HeaderValue::from_static("text/event-stream"))
314 .header(CONTENT_TYPE, HeaderValue::from_static("application/json"))
315 .body(payload)
316 .build()?;
317
318 if let Some(key) = self.api_key.as_deref() {
319 request.headers_mut().insert(
320 AUTHORIZATION,
321 HeaderValue::from_str(&format!("Bearer {key}")).unwrap(),
322 );
323 }
324
325 Ok(request)
326 }
327}
328
329#[cfg(test)]
330mod tests {
331 use super::*;
332 use meilisearch_test_macro::meilisearch_test;
333 use serde_json::json;
334 #[meilisearch_test]
335 async fn chat_workspace_lifecycle(client: Client, name: String) -> Result<(), Error> {
336 let _: serde_json::Value = client
337 .http_client
338 .request(
339 &format!("{}/experimental-features", client.host),
340 Method::Patch {
341 query: (),
342 body: &json!({ "chatCompletions": true }),
343 },
344 200,
345 )
346 .await?;
347
348 let workspace = format!("{name}-workspace");
349
350 let mut prompts = ChatPrompts::new();
351 prompts.set_system("You are a helpful assistant.");
352 prompts.set_search_description("Use search to fetch relevant documents.");
353
354 let mut settings = ChatWorkspaceSettings::new();
355 settings
356 .set_source("openAi")
357 .set_api_key("sk-test")
358 .set_prompts(prompts.clone());
359
360 let updated = client
361 .update_chat_workspace_settings(&workspace, &settings)
362 .await?;
363 assert_eq!(updated.source.as_deref(), Some("openAi"));
364 let updated_prompts = updated
365 .prompts
366 .expect("updated settings should contain prompts");
367 assert_eq!(updated_prompts.system.as_deref(), prompts.system.as_deref());
368 assert_eq!(
369 updated_prompts.search_description.as_deref(),
370 prompts.search_description.as_deref()
371 );
372 if let Some(masked_key) = updated.api_key.as_ref() {
373 assert_ne!(
374 masked_key, "sk-test",
375 "API key should not be returned in clear text"
376 );
377 }
378
379 let workspace_info = client.get_chat_workspace(&workspace).await?;
380 assert_eq!(workspace_info.uid, workspace);
381
382 let fetched_settings = client.get_chat_workspace_settings(&workspace).await?;
383 assert_eq!(fetched_settings.source.as_deref(), Some("openAi"));
384 let fetched_prompts = fetched_settings
385 .prompts
386 .expect("workspace should have prompts configured");
387 assert_eq!(fetched_prompts.system.as_deref(), prompts.system.as_deref());
388 assert_eq!(
389 fetched_prompts.search_description.as_deref(),
390 prompts.search_description.as_deref()
391 );
392
393 let list = client.list_chat_workspaces().await?;
394 assert!(list.results.iter().any(|w| w.uid == workspace));
395
396 let mut query = ChatWorkspacesQuery::new(&client);
397 query.with_limit(1);
398 let limited = query.execute().await?;
399 assert_eq!(limited.limit, 1);
400
401 let _ = client.reset_chat_workspace_settings(&workspace).await?;
402
403 Ok(())
404 }
405
406 #[test]
407 fn chat_prompts_builder_helpers() {
408 let mut prompts = ChatPrompts::new();
409 prompts
410 .set_system("system")
411 .set_search_description("desc")
412 .set_search_q_param("q")
413 .set_search_index_uid_param("idx")
414 .insert("custom", "value");
415
416 assert_eq!(prompts.system.as_deref(), Some("system"));
417 assert_eq!(prompts.search_description.as_deref(), Some("desc"));
418 assert_eq!(prompts.search_q_param.as_deref(), Some("q"));
419 assert_eq!(prompts.search_index_uid_param.as_deref(), Some("idx"));
420 assert_eq!(
421 prompts.extra.get("custom").map(String::as_str),
422 Some("value")
423 );
424 }
425
426 #[test]
427 fn chat_workspace_settings_builder_helpers() {
428 let mut settings = ChatWorkspaceSettings::new();
429 settings
430 .set_source("openAi")
431 .set_org_id("org")
432 .set_project_id("project")
433 .set_api_version("2024-01-01")
434 .set_deployment_id("deployment")
435 .set_base_url("http://example.com")
436 .set_api_key("secret")
437 .set_prompts({
438 let mut prompts = ChatPrompts::new();
439 prompts.set_system("hi");
440 prompts
441 });
442
443 assert_eq!(settings.source.as_deref(), Some("openAi"));
444 assert_eq!(settings.org_id.as_deref(), Some("org"));
445 assert_eq!(settings.project_id.as_deref(), Some("project"));
446 assert_eq!(settings.api_version.as_deref(), Some("2024-01-01"));
447 assert_eq!(settings.deployment_id.as_deref(), Some("deployment"));
448 assert_eq!(settings.base_url.as_deref(), Some("http://example.com"));
449 assert_eq!(settings.api_key.as_deref(), Some("secret"));
450 assert_eq!(
451 settings.prompts.and_then(|p| p.system).as_deref(),
452 Some("hi")
453 );
454 }
455
456 #[test]
457 #[cfg(feature = "reqwest")]
458 fn stream_chat_completion_request_includes_expected_headers() {
459 use reqwest::header::{AUTHORIZATION, CONTENT_TYPE};
460
461 let client = Client::new("http://localhost:7700", Some("secret")).unwrap();
462 let body = json!({
463 "model": "gpt-3.5-turbo",
464 "messages": [{ "role": "user", "content": "Hello" }],
465 "stream": true
466 });
467
468 let request = client
469 .build_stream_chat_request("workspace", &body)
470 .expect("request should be built");
471
472 assert_eq!(request.method(), reqwest::Method::POST);
473 assert_eq!(
474 request.url().as_str(),
475 "http://localhost:7700/chats/workspace/chat/completions"
476 );
477
478 let headers = request.headers();
479 assert_eq!(
480 headers
481 .get(reqwest::header::ACCEPT)
482 .map(|h| h.to_str().unwrap()),
483 Some("text/event-stream")
484 );
485 assert_eq!(
486 headers.get(CONTENT_TYPE).map(|h| h.to_str().unwrap()),
487 Some("application/json")
488 );
489 assert_eq!(
490 headers.get(AUTHORIZATION).map(|h| h.to_str().unwrap()),
491 Some("Bearer secret")
492 );
493
494 let expected_body = body.to_string();
495 let request_body = request
496 .body()
497 .and_then(|b| b.as_bytes())
498 .expect("request has body");
499 assert_eq!(request_body, expected_body.as_bytes());
500 }
501}