Skip to main content

meilisearch_sdk/
chats.rs

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/// Representation of a chat workspace.
13#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
14#[serde(rename_all = "camelCase")]
15pub struct ChatWorkspace {
16    pub uid: String,
17}
18
19/// Paginated chat workspace results.
20#[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/// Chat workspace prompts payload.
30#[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    /// Any additional provider-specific prompt values.
45    #[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/// Chat workspace settings payload.
82#[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/// Query builder for listing chat workspaces.
151#[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    /// List all chat workspaces.
188    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    /// List chat workspaces using query parameters.
199    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    /// Retrieve a chat workspace by uid.
213    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    /// Retrieve chat workspace settings.
224    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    /// Create or update chat workspace settings.
238    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    /// Reset chat workspace settings to defaults.
256    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    /// Stream chat completions for a workspace.
273    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}