Skip to main content

headless_lms_chatbot/chatbot_tools/provider_tools/
azure_ai_search.rs

1use crate::{
2    azure_chatbot::CONTENT_FIELD_SEPARATOR,
3    chatbot_error::chatbot_err,
4    prelude::{ChatbotError, ChatbotErrorType, ChatbotResult},
5    search_filter::SearchFilter,
6};
7use headless_lms_base::config::ApplicationConfiguration;
8use headless_lms_base::prelude_base_and_re_exports::BackendError;
9use serde::{Deserialize, Serialize};
10use url::Url;
11use uuid::Uuid;
12
13#[derive(Serialize, Deserialize, Debug, Clone)]
14pub struct AzureAISearchToolDefinition {
15    #[serde(rename = "type")]
16    pub data_type: String,
17    pub azure_ai_search: AzureAISearch,
18}
19
20#[derive(Serialize, Deserialize, Debug, Clone)]
21pub struct AzureAISearch {
22    pub indexes: Vec<SearchIndex>,
23}
24
25#[derive(Serialize, Deserialize, Debug, Clone)]
26pub struct SearchIndex {
27    pub project_connection_id: String,
28    pub index_name: String,
29    pub query_type: String,
30    pub top_k: i32,
31    pub embedding_dependency: EmbeddingDependency,
32    pub in_scope: bool,
33    pub strictness: i32,
34    #[serde(skip_serializing_if = "Option::is_none")]
35    pub filter: Option<String>,
36    pub fields_mapping: FieldsMapping,
37    pub semantic_configuration: String,
38}
39
40#[derive(Serialize, Deserialize, Debug, Clone)]
41pub struct FieldsMapping {
42    pub content_fields_separator: String,
43    pub content_fields: Vec<String>,
44    pub filepath_field: String,
45    pub title_field: String,
46    pub url_field: String,
47    pub vector_fields: Vec<String>,
48}
49
50#[derive(Serialize, Deserialize, Debug, Clone)]
51pub struct EmbeddingDependency {
52    #[serde(rename = "type")]
53    pub dep_type: String,
54    pub deployment_name: String,
55}
56
57pub fn get_azure_ai_search_tool_definition(
58    app_config: &ApplicationConfiguration,
59    course_id: Uuid,
60    use_semantic_reranking: bool,
61) -> ChatbotResult<AzureAISearchToolDefinition> {
62    let index_name = Url::parse(&app_config.base_url)?
63        .host_str()
64        .ok_or_else(|| {
65            chatbot_err!(
66                AzureRequestBuildError,
67                "Invalid application base url, no host"
68            )
69        })?
70        .replace(".", "-");
71    let azure_config = app_config.azure_configuration.as_ref().ok_or_else(|| {
72        chatbot_err!(
73            AzureRequestBuildError,
74            "Azure configuration is missing from the application configuration"
75        )
76    })?;
77
78    let search_config = azure_config.search_config.as_ref().ok_or_else(|| {
79        chatbot_err!(
80            AzureRequestBuildError,
81            "Azure search configuration is missing from the Azure configuration"
82        )
83    })?;
84
85    let query_type = if use_semantic_reranking {
86        "vector_semantic_hybrid"
87    } else {
88        "vector_simple_hybrid"
89    };
90
91    let semantic_configuration = format!("{}-semantic-configuration", &index_name);
92
93    Ok(AzureAISearchToolDefinition {
94        data_type: "azure_ai_search".to_string(),
95        azure_ai_search: AzureAISearch {
96            indexes: vec![SearchIndex {
97                index_name,
98                project_connection_id: search_config.search_connection_id.to_owned(),
99                query_type: query_type.to_string(),
100                semantic_configuration,
101                embedding_dependency: EmbeddingDependency {
102                    dep_type: "deployment_name".to_string(),
103                    deployment_name: search_config.vectorizer_deployment_id.clone(),
104                },
105                in_scope: false,
106                top_k: 15,
107                strictness: 3,
108                filter: Some(SearchFilter::eq("course_id", course_id.to_string()).to_odata()?),
109                fields_mapping: FieldsMapping {
110                    content_fields_separator: CONTENT_FIELD_SEPARATOR.to_string(),
111                    content_fields: vec!["chunk_context".to_string(), "chunk".to_string()],
112                    filepath_field: "filepath".to_string(),
113                    title_field: "title".to_string(),
114                    url_field: "url".to_string(),
115                    vector_fields: vec!["text_vector".to_string()],
116                },
117            }],
118        },
119    })
120}