headless_lms_chatbot/chatbot_tools/custom_tools/
document_lookup.rs1use headless_lms_utils::cache::Cache;
2use indexmap::IndexMap;
3
4use headless_lms_models::chatbot_configurations::ToolCategory;
5use headless_lms_models::{course_page_markdown_content, pages};
6use headless_lms_utils::{
7 document_schema_processor::remove_sensitive_attributes,
8 json_schema_types::{JSONType, JsonItem, Schema, SchemaPropertyType},
9 strings::truncate_utf8_at_boundary,
10};
11
12use crate::{
13 azure_chatbot::azure::tools::{AzureLLMFunctionToolDefinition, LLMToolType},
14 chatbot_tools::{
15 ChatbotTool, ChatbotToolDeclaration, ToolProperties,
16 argument_parsing::deserialize_to_optional_uuid_and_errors_to_none,
17 course_scope::{
18 COURSE_ID_ARGUMENT_DESCRIPTION, material_requirements, resolve_course_scope,
19 },
20 tool_authorization::ToolRequirement,
21 },
22 citations::parse_document_filepath,
23 llm_utils::estimate_tokens,
24 prelude::*,
25 user_context::ChatbotTurnContext,
26};
27
28pub type DocumentLookupTool = ToolProperties<DocumentLookupState>;
29
30pub struct DocumentLookupState {
31 document: Option<String>,
32}
33
34#[derive(Deserialize)]
35pub struct DocumentLookupArguments {
36 #[allow(dead_code)]
39 title: String,
40 filepath: Option<String>,
41 #[serde(deserialize_with = "deserialize_to_optional_uuid_and_errors_to_none")]
42 page_id: Option<Uuid>,
43 format: String,
44 #[serde(deserialize_with = "deserialize_to_optional_uuid_and_errors_to_none")]
45 course_id: Option<Uuid>,
46}
47
48fn shorten_page_content(mut content: String) -> String {
52 const MAX_TOKENS: i32 = 25_000;
53 loop {
54 let tokens = estimate_tokens(&content);
55 if tokens <= MAX_TOKENS {
56 return content;
57 }
58 let max_bytes = content.len() * (MAX_TOKENS as usize - 1_000) / tokens as usize;
59 content = truncate_utf8_at_boundary(&content, max_bytes).to_string();
60 }
61}
62
63impl ChatbotToolDeclaration for DocumentLookupTool {
64 const NAME: &'static str = "document_lookup";
65
66 fn offer_requirements(user_context: &ChatbotTurnContext) -> Vec<ToolRequirement> {
67 material_requirements(user_context.course_id)
68 }
69
70 const CATEGORY: ToolCategory = ToolCategory::CourseMaterial;
71
72 fn get_tool_definition() -> AzureLLMFunctionToolDefinition {
73 AzureLLMFunctionToolDefinition {
74 tool_type: LLMToolType::Function,
75 name: Self::NAME.to_string(),
76 description: "Look up the full content of a specific document by the title and filepath or id (page_id). The needed arguments can be found from Azure search results or by using the course_structure tool. Either a filepath or a page_id is required to find the correct document, in addition to the document title. The document can be returned in Markdown or JSON format. The Markdown format is cleaner and preferred, but might have errors: if you suspect it's erroneous, you can request the JSON version.".to_string(),
77 parameters: Schema::strict_object(
78 IndexMap::from([
79 (
80 "filepath".to_string(),
81 SchemaPropertyType::Item(JsonItem {
82 type_field: JSONType::String,
83 description: Some("The filepath of the document to look up, as returned from Azure search. Either the filepath or page_id is required.".to_string()),
84 }),
85 ),
86 (
87 "title".to_string(),
88 SchemaPropertyType::Item(JsonItem {
89 type_field: JSONType::String,
90 description: Some("The title of the document to look up, as returned from Azure search. Optional.".to_string()),
91 }),
92 ),
93 (
94 "page_id".to_string(),
95 SchemaPropertyType::Item(JsonItem {
96 type_field: JSONType::String,
97 description: Some("The page_id of the document to look up. Either page_id or the filepath is required.".to_string()),
98 }),
99 ),
100 (
101 "format".to_string(),
102 SchemaPropertyType::Item(JsonItem {
103 type_field: JSONType::String,
104 description: Some("The format of the document. Optional. Valid values are 'json' and 'markdown'. Markdown content is human readable, but might have errors. ".to_string()),
105 }),
106 ),
107 (
108 "course_id".to_string(),
109 SchemaPropertyType::Item(JsonItem {
110 type_field: JSONType::String,
111 description: Some(COURSE_ID_ARGUMENT_DESCRIPTION.to_string()),
112 }),
113 )
114 ]),
115 None,
116 ),
117 strict: true,
118 }
119 }
120}
121
122impl ChatbotTool for DocumentLookupTool {
124 type Arguments = DocumentLookupArguments;
125
126 fn call_requirements(
127 arguments: &Self::Arguments,
128 user_context: &ChatbotTurnContext,
129 ) -> Vec<ToolRequirement> {
130 material_requirements(resolve_course_scope(user_context, arguments.course_id).ok())
131 }
132
133 async fn from_db_and_arguments(
134 conn: &mut PgConnection,
135 _app_config: &ApplicationConfiguration,
136 _cache: &Cache,
137 arguments: Self::Arguments,
138 user_context: &ChatbotTurnContext,
139 ) -> ChatbotResult<Self> {
140 let course_id = resolve_course_scope(user_context, arguments.course_id)?;
141
142 let page_id = if let Some(id) = &arguments.page_id {
143 id.to_owned()
144 } else if let Some(f) = &arguments.filepath {
145 let res = parse_document_filepath(f);
146 match res {
147 Ok(d) => d.page_id,
148 Err(e) => Err(chatbot_err!(
149 InvalidToolArguments,
150 "Couldn't parse document file path and no valid page id was provided, unable to look up document.",
151 e
152 ))?,
153 }
154 } else {
155 return Err(chatbot_err!(
156 InvalidToolArguments,
157 format!(
158 "Unable to call document_lookup tool. No filepath or page id provided. One of them is needed to find the document."
159 )
160 ));
161 };
162 let document = match course_page_markdown_content::get_course_page_content_by_page_id(
163 conn, page_id,
164 )
165 .await
166 {
167 Ok(page_content) if page_content.course_id == course_id => {
169 if arguments.format == "json" {
170 let s =
171 shorten_page_content(serde_json::to_string(&page_content.json_content)?);
172 Some(s)
173 } else if let Some(content) = page_content.markdown_content {
174 let s = shorten_page_content(content);
175 Some(s)
176 } else {
177 let base = "Markdown content not found. Page JSON content:\n\n".to_string();
178 let s =
179 shorten_page_content(serde_json::to_string(&page_content.json_content)?);
180 Some(base + &s)
181 }
182 }
183 Ok(_) => None,
184 Err(e) if e.error_type() == &ModelErrorType::RecordNotFound => {
188 match pages::get_page(conn, page_id).await {
189 Ok(page) if page.course_id == Some(course_id) && page.deleted_at.is_none() => {
190 let blocks = remove_sensitive_attributes(page.blocks_cloned()?);
191 let base = "No converted markdown exists for this course; this is raw block JSON:\n\n".to_string();
192 let s = shorten_page_content(serde_json::to_string(&blocks)?);
193 Some(base + &s)
194 }
195 _ => None,
196 }
197 }
198 Err(e) => return Err(ChatbotError::from(e)),
199 };
200
201 Ok(DocumentLookupTool {
202 state: DocumentLookupState { document },
203 })
204 }
205
206 fn output(&self) -> String {
207 if let Some(d) = &self.state.document {
208 d.to_string()
209 } else {
210 "Document not found.".to_string()
211 }
212 }
213
214 fn output_description_instructions(&self) -> Option<String> {
215 Some("Do not return the whole document to the user. Use the document as a source of more information for answering the user etc. Cite the course_material_search result the page came from; document_lookup itself produces no citation.".to_string())
216 }
217}
218
219#[cfg(test)]
220mod tests {
221 use super::*;
222
223 #[test]
224 fn shorten_page_content_shortens_prose() {
225 let input = "The quick brown fox jumps over the lazy dog. ".repeat(4600);
226 assert!(input.len() > 200_000);
227
228 let shortened = shorten_page_content(input);
229
230 assert!(estimate_tokens(&shortened) <= 25_000);
231 }
232
233 #[test]
234 fn shorten_page_content_shortens_punctuation_heavy_json() {
235 let input = format!(
236 "[{}]",
237 r#"{"id":"1","name":"block","attributes":{"content":"Hei, mitä kuuluu?"}},"#
238 .repeat(2500)
239 );
240 assert!(input.len() > 150_000);
241
242 let shortened = shorten_page_content(input);
243
244 assert!(estimate_tokens(&shortened) <= 25_000);
245 }
246
247 #[test]
248 fn shorten_page_content_leaves_short_content_alone() {
249 let input = "Short enough.".to_string();
250
251 assert_eq!(shorten_page_content(input.clone()), input);
252 }
253}