1use std::collections::HashMap;
18
19use serde::Serialize;
20use validator::*;
21
22use super::model_validate::validate_json_schema_value;
23use crate::tool::web_search::request::{ContentSize, SearchEngine, SearchRecencyFilter};
24
25#[derive(Debug, Clone, Serialize)]
55pub struct ThinkingType {
56 #[serde(rename = "type")]
58 pub mode: ThinkingMode,
59
60 #[serde(skip_serializing_if = "Option::is_none")]
66 pub clear_thinking: Option<bool>,
67}
68
69#[derive(Debug, Clone, Serialize)]
71#[serde(rename_all = "lowercase")]
72pub enum ThinkingMode {
73 Enabled,
74 Disabled,
75}
76
77impl ThinkingType {
78 pub fn enabled() -> Self {
80 Self {
81 mode: ThinkingMode::Enabled,
82 clear_thinking: None,
83 }
84 }
85
86 pub fn disabled() -> Self {
88 Self {
89 mode: ThinkingMode::Disabled,
90 clear_thinking: None,
91 }
92 }
93
94 pub fn with_clear_thinking(mut self, clear: bool) -> Self {
99 self.clear_thinking = Some(clear);
100 self
101 }
102}
103
104#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
134pub enum ReasoningEffort {
135 #[serde(rename = "max")]
138 Max,
139 #[serde(rename = "xhigh")]
141 Xhigh,
142 #[serde(rename = "high")]
144 High,
145 #[serde(rename = "medium")]
147 Medium,
148 #[serde(rename = "low")]
150 Low,
151 #[serde(rename = "minimal")]
153 Minimal,
154 #[serde(rename = "none")]
156 None,
157}
158
159#[derive(Debug, Clone, Serialize)]
200#[serde(tag = "type")]
201#[serde(rename_all = "snake_case")]
202pub enum Tools {
203 Function { function: Function },
209
210 Retrieval { retrieval: Retrieval },
215
216 WebSearch { web_search: WebSearch },
221
222 #[serde(rename = "mcp")]
227 MCP { mcp: MCP },
228}
229
230#[derive(Debug, Clone, Serialize, Validate)]
240pub struct Function {
241 #[validate(length(min = 1, max = 64))]
243 pub name: String,
244
245 pub description: String,
247
248 #[serde(skip_serializing_if = "Option::is_none")]
252 #[validate(custom(function = "validate_json_schema_value"))]
253 pub parameters: Option<serde_json::Value>,
254}
255
256impl Function {
257 pub fn new(
279 name: impl Into<String>,
280 description: impl Into<String>,
281 parameters: serde_json::Value,
282 ) -> Self {
283 Self {
284 name: name.into(),
285 description: description.into(),
286 parameters: Some(parameters),
287 }
288 }
289}
290
291#[derive(Debug, Clone, Serialize)]
296pub struct Retrieval {
297 knowledge_id: String,
298 #[serde(skip_serializing_if = "Option::is_none")]
299 prompt_template: Option<String>,
300}
301
302impl Retrieval {
303 pub fn new(knowledge_id: impl Into<String>, prompt_template: Option<String>) -> Self {
305 Self {
306 knowledge_id: knowledge_id.into(),
307 prompt_template,
308 }
309 }
310}
311
312#[derive(Debug, Clone, Serialize, PartialEq)]
316#[serde(rename_all = "snake_case")]
317pub enum ResultSequence {
318 Before,
319 After,
320}
321
322#[derive(Debug, Clone, Serialize, Validate)]
325pub struct WebSearch {
326 pub search_engine: SearchEngine,
329
330 #[serde(skip_serializing_if = "Option::is_none")]
332 pub enable: Option<bool>,
333
334 #[serde(skip_serializing_if = "Option::is_none")]
336 pub search_query: Option<String>,
337
338 #[serde(skip_serializing_if = "Option::is_none")]
341 pub search_intent: Option<bool>,
342
343 #[serde(skip_serializing_if = "Option::is_none")]
345 #[validate(range(min = 1, max = 50))]
346 pub count: Option<u32>,
347
348 #[serde(skip_serializing_if = "Option::is_none")]
350 pub search_domain_filter: Option<String>,
351
352 #[serde(skip_serializing_if = "Option::is_none")]
354 pub search_recency_filter: Option<SearchRecencyFilter>,
355
356 #[serde(skip_serializing_if = "Option::is_none")]
358 pub content_size: Option<ContentSize>,
359
360 #[serde(skip_serializing_if = "Option::is_none")]
362 pub result_sequence: Option<ResultSequence>,
363
364 #[serde(skip_serializing_if = "Option::is_none")]
366 pub search_result: Option<bool>,
367
368 #[serde(skip_serializing_if = "Option::is_none")]
370 pub require_search: Option<bool>,
371
372 #[serde(skip_serializing_if = "Option::is_none")]
374 pub search_prompt: Option<String>,
375}
376
377impl WebSearch {
378 pub fn new(search_engine: SearchEngine) -> Self {
381 Self {
382 search_engine,
383 enable: None,
384 search_query: None,
385 search_intent: None,
386 count: None,
387 search_domain_filter: None,
388 search_recency_filter: None,
389 content_size: None,
390 result_sequence: None,
391 search_result: None,
392 require_search: None,
393 search_prompt: None,
394 }
395 }
396
397 pub fn with_enable(mut self, enable: bool) -> Self {
399 self.enable = Some(enable);
400 self
401 }
402 pub fn with_search_query(mut self, query: impl Into<String>) -> Self {
404 self.search_query = Some(query.into());
405 self
406 }
407 pub fn with_search_intent(mut self, search_intent: bool) -> Self {
409 self.search_intent = Some(search_intent);
410 self
411 }
412 pub fn with_count(mut self, count: u32) -> Self {
414 self.count = Some(count);
415 self
416 }
417 pub fn with_search_domain_filter(mut self, domain: impl Into<String>) -> Self {
419 self.search_domain_filter = Some(domain.into());
420 self
421 }
422 pub fn with_search_recency_filter(mut self, filter: SearchRecencyFilter) -> Self {
424 self.search_recency_filter = Some(filter);
425 self
426 }
427 pub fn with_content_size(mut self, size: ContentSize) -> Self {
429 self.content_size = Some(size);
430 self
431 }
432 pub fn with_result_sequence(mut self, seq: ResultSequence) -> Self {
434 self.result_sequence = Some(seq);
435 self
436 }
437 pub fn with_search_result(mut self, enable: bool) -> Self {
439 self.search_result = Some(enable);
440 self
441 }
442 pub fn with_require_search(mut self, require: bool) -> Self {
444 self.require_search = Some(require);
445 self
446 }
447 pub fn with_search_prompt(mut self, prompt: impl Into<String>) -> Self {
449 self.search_prompt = Some(prompt.into());
450 self
451 }
452}
453#[derive(Debug, Clone, Serialize, Validate)]
457pub struct MCP {
458 #[validate(length(min = 1))]
461 pub server_label: String,
462
463 #[serde(skip_serializing_if = "Option::is_none")]
465 #[validate(url)]
466 pub server_url: Option<String>,
467
468 #[serde(skip_serializing_if = "Option::is_none")]
470 pub transport_type: Option<MCPTransportType>,
471
472 #[serde(skip_serializing_if = "Vec::is_empty")]
474 pub allowed_tools: Vec<String>,
475
476 #[serde(skip_serializing_if = "Option::is_none")]
478 pub headers: Option<HashMap<String, String>>,
479}
480
481impl MCP {
482 pub fn new(server_label: impl Into<String>) -> Self {
485 Self {
486 server_label: server_label.into(),
487 server_url: None,
488 transport_type: Some(MCPTransportType::StreamableHttp),
489 allowed_tools: Vec::new(),
490 headers: None,
491 }
492 }
493
494 pub fn with_server_url(mut self, url: impl Into<String>) -> Self {
496 self.server_url = Some(url.into());
497 self
498 }
499 pub fn with_transport_type(mut self, transport: MCPTransportType) -> Self {
501 self.transport_type = Some(transport);
502 self
503 }
504 pub fn with_allowed_tools(mut self, tools: impl Into<Vec<String>>) -> Self {
506 self.allowed_tools = tools.into();
507 self
508 }
509 pub fn add_allowed_tool(mut self, tool: impl Into<String>) -> Self {
511 self.allowed_tools.push(tool.into());
512 self
513 }
514 pub fn with_headers(mut self, headers: HashMap<String, String>) -> Self {
516 self.headers = Some(headers);
517 self
518 }
519 pub fn with_header(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
521 let mut map = self.headers.unwrap_or_default();
522 map.insert(key.into(), value.into());
523 self.headers = Some(map);
524 self
525 }
526}
527
528#[derive(Debug, Clone, Serialize, PartialEq)]
530#[serde(rename_all = "kebab-case")]
531pub enum MCPTransportType {
532 Sse,
533 StreamableHttp,
534}
535
536#[derive(Debug, Clone, Copy, Serialize)]
546#[serde(rename_all = "snake_case")]
547#[serde(tag = "type")]
548pub enum ResponseFormat {
549 Text,
551 JsonObject,
553}
554
555#[cfg(test)]
556mod tests {
557 use super::*;
558
559 #[test]
561 fn test_thinking_type_enabled_serialization() {
562 let thinking = ThinkingType::enabled();
563 let json = serde_json::to_string(&thinking).unwrap();
564 assert!(json.contains("\"type\":\"enabled\""));
565 assert!(!json.contains("clear_thinking"));
566 }
567
568 #[test]
569 fn test_thinking_type_disabled_serialization() {
570 let thinking = ThinkingType::disabled();
571 let json = serde_json::to_string(&thinking).unwrap();
572 assert!(json.contains("\"type\":\"disabled\""));
573 assert!(!json.contains("clear_thinking"));
574 }
575
576 #[test]
577 fn test_thinking_type_with_clear_thinking_serialization() {
578 let thinking = ThinkingType::enabled().with_clear_thinking(false);
579 let json = serde_json::to_string(&thinking).unwrap();
580 assert!(json.contains("\"type\":\"enabled\""));
581 assert!(json.contains("\"clear_thinking\":false"));
582 }
583
584 #[test]
585 fn test_thinking_type_disabled_with_clear_thinking() {
586 let thinking = ThinkingType::disabled().with_clear_thinking(true);
587 let json = serde_json::to_string(&thinking).unwrap();
588 assert!(json.contains("\"type\":\"disabled\""));
589 assert!(json.contains("\"clear_thinking\":true"));
590 }
591
592 #[test]
594 fn test_function_new() {
595 let params = serde_json::json!({
596 "type": "object",
597 "properties": {
598 "name": {"type": "string"}
599 }
600 });
601 let func = Function::new("test_func", "A test function", params);
602
603 assert_eq!(func.name, "test_func");
604 assert_eq!(func.description, "A test function");
605 assert!(func.parameters.is_some());
606 }
607
608 #[test]
609 fn test_function_serialization() {
610 let params = serde_json::json!({
611 "type": "object",
612 "properties": {
613 "value": {"type": "number"}
614 }
615 });
616 let func = Function::new("test_func", "A test function", params);
617 let json = serde_json::to_string(&func).unwrap();
618
619 assert!(json.contains("\"name\":\"test_func\""));
620 assert!(json.contains("\"description\":\"A test function\""));
621 assert!(json.contains("\"properties\""));
622 }
623
624 #[test]
625 fn test_function_validation() {
626 let params = serde_json::json!({
627 "type": "object",
628 "properties": {}
629 });
630 let func = Function::new("valid_name", "Description", params.clone());
631
632 assert!(func.validate().is_ok());
634
635 let invalid_name = Function::new("", "Description", params.clone());
636 assert!(invalid_name.validate().is_err());
637
638 let long_name = Function::new("a".repeat(65), "Description", params);
639 assert!(long_name.validate().is_err());
640 }
641
642 #[test]
644 fn test_retrieval_new() {
645 let retrieval = Retrieval::new("kb_123", Some("template".to_string()));
646 assert_eq!(retrieval.knowledge_id, "kb_123");
647 assert_eq!(retrieval.prompt_template, Some("template".to_string()));
648 }
649
650 #[test]
651 fn test_retrieval_new_without_template() {
652 let retrieval = Retrieval::new("kb_456", None);
653 assert_eq!(retrieval.knowledge_id, "kb_456");
654 assert!(retrieval.prompt_template.is_none());
655 }
656
657 #[test]
658 fn test_retrieval_serialization() {
659 let retrieval = Retrieval::new("kb_789", None);
660 let json = serde_json::to_string(&retrieval).unwrap();
661 assert!(json.contains("\"knowledge_id\":\"kb_789\""));
662 assert!(!json.contains("prompt_template"));
664 }
665
666 #[test]
668 fn test_web_search_new() {
669 let web_search = WebSearch::new(SearchEngine::SearchPro);
670 assert_eq!(web_search.search_engine, SearchEngine::SearchPro);
671 assert!(web_search.enable.is_none());
672 }
673
674 #[test]
675 fn test_web_search_with_enable() {
676 let web_search = WebSearch::new(SearchEngine::SearchPro).with_enable(true);
677 assert_eq!(web_search.enable, Some(true));
678 }
679
680 #[test]
681 fn test_web_search_with_search_query() {
682 let web_search = WebSearch::new(SearchEngine::SearchPro).with_search_query("test query");
683 assert_eq!(web_search.search_query, Some("test query".to_string()));
684 }
685
686 #[test]
687 fn test_web_search_with_search_intent() {
688 let web_search = WebSearch::new(SearchEngine::SearchPro).with_search_intent(true);
689 assert_eq!(web_search.search_intent, Some(true));
690 }
691
692 #[test]
693 fn test_web_search_with_count() {
694 let web_search = WebSearch::new(SearchEngine::SearchPro).with_count(10);
695 assert_eq!(web_search.count, Some(10));
696 }
697
698 #[test]
699 fn test_web_search_with_search_domain_filter() {
700 let web_search =
701 WebSearch::new(SearchEngine::SearchPro).with_search_domain_filter("example.com");
702 assert_eq!(
703 web_search.search_domain_filter,
704 Some("example.com".to_string())
705 );
706 }
707
708 #[test]
709 fn test_web_search_with_search_recency_filter() {
710 let filter = SearchRecencyFilter::OneDay;
711 let web_search =
712 WebSearch::new(SearchEngine::SearchPro).with_search_recency_filter(filter.clone());
713 assert_eq!(web_search.search_recency_filter, Some(filter));
714 }
715
716 #[test]
717 fn test_web_search_with_content_size() {
718 let size = ContentSize::Medium;
719 let web_search = WebSearch::new(SearchEngine::SearchPro).with_content_size(size.clone());
720 assert_eq!(web_search.content_size, Some(size));
721 }
722
723 #[test]
724 fn test_web_search_with_result_sequence() {
725 let seq = ResultSequence::After;
726 let web_search = WebSearch::new(SearchEngine::SearchPro).with_result_sequence(seq.clone());
727 assert_eq!(web_search.result_sequence, Some(seq));
728 }
729
730 #[test]
731 fn test_web_search_with_search_result() {
732 let web_search = WebSearch::new(SearchEngine::SearchPro).with_search_result(true);
733 assert_eq!(web_search.search_result, Some(true));
734 }
735
736 #[test]
737 fn test_web_search_with_require_search() {
738 let web_search = WebSearch::new(SearchEngine::SearchPro).with_require_search(true);
739 assert_eq!(web_search.require_search, Some(true));
740 }
741
742 #[test]
743 fn test_web_search_with_search_prompt() {
744 let web_search =
745 WebSearch::new(SearchEngine::SearchPro).with_search_prompt("custom prompt");
746 assert_eq!(web_search.search_prompt, Some("custom prompt".to_string()));
747 }
748
749 #[test]
750 fn test_web_search_serialization() {
751 let web_search = WebSearch::new(SearchEngine::SearchPro)
752 .with_enable(true)
753 .with_count(5);
754 let json = serde_json::to_string(&web_search).unwrap();
755 assert!(json.contains("\"search_engine\""));
756 assert!(json.contains("\"enable\":true"));
757 assert!(json.contains("\"count\":5"));
758 }
759
760 #[test]
762 fn test_mcp_new() {
763 let mcp = MCP::new("server_label");
764 assert_eq!(mcp.server_label, "server_label");
765 assert_eq!(mcp.transport_type, Some(MCPTransportType::StreamableHttp));
766 assert!(mcp.allowed_tools.is_empty());
767 }
768
769 #[test]
770 fn test_mcp_with_server_url() {
771 let mcp = MCP::new("server_label").with_server_url("https://example.com");
772 assert_eq!(mcp.server_url, Some("https://example.com".to_string()));
773 }
774
775 #[test]
776 fn test_mcp_with_transport_type() {
777 let mcp = MCP::new("server_label").with_transport_type(MCPTransportType::Sse);
778 assert_eq!(mcp.transport_type, Some(MCPTransportType::Sse));
779 }
780
781 #[test]
782 fn test_mcp_with_allowed_tools() {
783 let mcp = MCP::new("server_label")
784 .with_allowed_tools(vec!["tool1".to_string(), "tool2".to_string()]);
785 assert_eq!(mcp.allowed_tools.len(), 2);
786 assert!(mcp.allowed_tools.contains(&"tool1".to_string()));
787 }
788
789 #[test]
790 fn test_mcp_add_allowed_tool() {
791 let mcp = MCP::new("server_label")
792 .add_allowed_tool("tool1")
793 .add_allowed_tool("tool2");
794 assert_eq!(mcp.allowed_tools.len(), 2);
795 }
796
797 #[test]
798 fn test_mcp_with_headers() {
799 let mut headers = HashMap::new();
800 headers.insert("Authorization".to_string(), "Bearer token".to_string());
801 let mcp = MCP::new("server_label").with_headers(headers.clone());
802 assert_eq!(mcp.headers, Some(headers));
803 }
804
805 #[test]
806 fn test_mcp_with_header() {
807 let mcp = MCP::new("server_label").with_header("Authorization", "Bearer token");
808 let headers = mcp.headers.unwrap();
809 assert_eq!(
810 headers.get("Authorization"),
811 Some(&"Bearer token".to_string())
812 );
813 }
814
815 #[test]
816 fn test_mcp_serialization() {
817 let mcp = MCP::new("server_label")
818 .with_server_url("https://example.com")
819 .with_transport_type(MCPTransportType::Sse);
820 let json = serde_json::to_string(&mcp).unwrap();
821 assert!(json.contains("\"server_label\":\"server_label\""));
822 assert!(json.contains("\"server_url\":\"https://example.com\""));
823 assert!(json.contains("\"transport_type\":\"sse\""));
824 assert!(!json.contains("allowed_tools"));
826 }
827
828 #[test]
830 fn test_mcp_transport_type_sse_serialization() {
831 let transport = MCPTransportType::Sse;
832 let json = serde_json::to_string(&transport).unwrap();
833 assert!(json.contains("\"sse\""));
834 }
835
836 #[test]
837 fn test_mcp_transport_type_streamable_http_serialization() {
838 let transport = MCPTransportType::StreamableHttp;
839 let json = serde_json::to_string(&transport).unwrap();
840 assert!(json.contains("\"streamable-http\""));
841 }
842
843 #[test]
845 fn test_response_format_text_serialization() {
846 let format = ResponseFormat::Text;
847 let json = serde_json::to_string(&format).unwrap();
848 assert!(json.contains("\"type\":\"text\""));
849 }
850
851 #[test]
852 fn test_response_format_json_object_serialization() {
853 let format = ResponseFormat::JsonObject;
854 let json = serde_json::to_string(&format).unwrap();
855 assert!(json.contains("\"type\":\"json_object\""));
856 }
857
858 #[test]
860 fn test_tools_function_serialization() {
861 let func = Function::new("test_func", "test", serde_json::json!({}));
862 let tools = Tools::Function { function: func };
863 let json = serde_json::to_string(&tools).unwrap();
864 assert!(json.contains("\"type\":\"function\""));
865 assert!(json.contains("\"name\":\"test_func\""));
866 }
867
868 #[test]
869 fn test_tools_retrieval_serialization() {
870 let retrieval = Retrieval::new("kb_123", None);
871 let tools = Tools::Retrieval { retrieval };
872 let json = serde_json::to_string(&tools).unwrap();
873 assert!(json.contains("\"type\":\"retrieval\""));
874 assert!(json.contains("\"knowledge_id\":\"kb_123\""));
875 }
876
877 #[test]
878 fn test_tools_web_search_serialization() {
879 let web_search = WebSearch::new(SearchEngine::SearchPro);
880 let tools = Tools::WebSearch { web_search };
881 let json = serde_json::to_string(&tools).unwrap();
882 assert!(json.contains("\"type\":\"web_search\""));
883 assert!(json.contains("\"search_engine\""));
884 }
885
886 #[test]
887 fn test_tools_mcp_serialization() {
888 let mcp = MCP::new("server_label");
889 let tools = Tools::MCP { mcp };
890 let json = serde_json::to_string(&tools).unwrap();
891 assert!(json.contains("\"type\":\"mcp\""));
892 assert!(json.contains("\"server_label\":\"server_label\""));
893 }
894
895 #[test]
897 fn test_result_sequence_before_serialization() {
898 let seq = ResultSequence::Before;
899 let json = serde_json::to_string(&seq).unwrap();
900 assert!(json.contains("\"before\""));
901 }
902
903 #[test]
904 fn test_result_sequence_after_serialization() {
905 let seq = ResultSequence::After;
906 let json = serde_json::to_string(&seq).unwrap();
907 assert!(json.contains("\"after\""));
908 }
909}