1use std::time::Duration;
2
3use bytes::Bytes;
4use reqwest::multipart;
5use serde_json::Value;
6
7use crate::document::{InputDocument, InputKind, OutputFormat};
8use crate::error::{PdfConvertError, Result};
9use crate::models::vlm::{OpenRouterConfigBuilder, VlmConvertOptions};
10use crate::models::{
11 CodeFormulaVlmOptions, PictureDescriptionVlmEngineOptions, TaskPostResponse, TaskStatusResponse,
12};
13
14use super::client::{
15 extract_error_details, get_request, get_request_with_conn_close, handle_response,
16 retry_with_backoff,
17};
18
19#[derive(Debug, Clone)]
20pub struct DoclingConfig {
21 pub base_url: String,
22 pub openai_base_url: String,
23 pub vlm_pipeline_model: String,
24 pub picture_description_model: String,
25 pub code_formula_model: String,
26 pub api_key: Option<String>,
27}
28
29#[derive(Debug, Clone)]
30pub struct DoclingConvertRequest {
31 pub output_formats: Vec<OutputFormat>,
32 pub page_range: Option<(u32, u32)>,
33 pub chunking: bool,
34}
35
36impl DoclingConvertRequest {
37 pub fn for_outputs(output_formats: Vec<OutputFormat>) -> Self {
38 Self {
39 output_formats,
40 page_range: None,
41 chunking: false,
42 }
43 }
44}
45
46#[derive(Debug, Clone)]
47pub struct DoclingClient {
48 http_client: reqwest::Client,
49 config: DoclingConfig,
50}
51
52impl DoclingClient {
53 pub fn new(config: DoclingConfig) -> Result<Self> {
54 let http_client = reqwest::Client::builder()
55 .timeout(Duration::from_secs(300))
56 .tcp_keepalive(Duration::from_secs(60))
57 .pool_idle_timeout(Duration::from_secs(30))
58 .build()
59 .map_err(|e| PdfConvertError::api_error(None, e.to_string()))?;
60
61 Ok(Self {
62 http_client,
63 config,
64 })
65 }
66
67 pub fn config(&self) -> &DoclingConfig {
68 &self.config
69 }
70
71 pub async fn convert_file(
72 &self,
73 input: &InputDocument,
74 request: &DoclingConvertRequest,
75 ) -> Result<Value> {
76 let operation = || async {
77 let form = self.build_form(input, request)?;
78 let url = format!("{}/convert/file", self.config.base_url);
79 let response = self
80 .http_client
81 .post(&url)
82 .multipart(form)
83 .send()
84 .await
85 .map_err(PdfConvertError::from)?;
86
87 let response = handle_response(response, "Docling file conversion").await?;
88 response.json::<Value>().await.map_err(|e| {
89 PdfConvertError::parse_error("Docling conversion response", e.to_string())
90 })
91 };
92
93 retry_with_backoff(operation, "docling_convert_file").await
94 }
95
96 pub async fn submit_file_async(
97 &self,
98 input: &InputDocument,
99 request: &DoclingConvertRequest,
100 ) -> Result<String> {
101 let operation = || async {
102 let form = self.build_form(input, request)?;
103 let url = format!("{}/convert/file/async", self.config.base_url);
104 let response = self
105 .http_client
106 .post(&url)
107 .multipart(form)
108 .send()
109 .await
110 .map_err(PdfConvertError::from)?;
111
112 let response = handle_response(response, "Docling async submission").await?;
113 let result = response.json::<TaskPostResponse>().await.map_err(|e| {
114 PdfConvertError::parse_error("Docling async submission response", e.to_string())
115 })?;
116 Ok(result.task_id)
117 };
118
119 retry_with_backoff(operation, "docling_submit_file_async").await
120 }
121
122 pub async fn wait_for_result(&self, task_id: &str) -> Result<Value> {
123 loop {
124 match self.check_task_status(task_id).await? {
125 true => {
126 tokio::time::sleep(Duration::from_millis(500)).await;
127 return self.get_task_result(task_id, true).await;
128 }
129 false => tokio::time::sleep(Duration::from_secs(5)).await,
130 }
131 }
132 }
133
134 pub async fn check_task_status(&self, task_id: &str) -> Result<bool> {
135 let operation = || async {
136 let url = format!("{}/status/poll/{}", self.config.base_url, task_id);
137 let response = get_request(&self.http_client, &url, "Polling task status").await?;
138 let response_text = response.text().await?;
139
140 match serde_json::from_str::<TaskStatusResponse>(&response_text) {
141 Ok(status_response) => match status_response.task_status.as_str() {
142 "success" => Ok(true),
143 "failure" | "revoked" => {
144 let error_details = extract_error_details(&response_text);
145 Err(PdfConvertError::api_task_failed(
146 status_response.task_status,
147 error_details,
148 ))
149 }
150 _ => Ok(false),
151 },
152 Err(_) => Err(PdfConvertError::parse_error(
153 "task status response",
154 format!(
155 "Task {} returned invalid response: {}",
156 task_id, response_text
157 ),
158 )),
159 }
160 };
161
162 retry_with_backoff(operation, &format!("check_task_status({task_id})")).await
163 }
164
165 pub async fn get_task_result(&self, task_id: &str, use_new_conn: bool) -> Result<Value> {
166 let operation = || async {
167 let url = format!("{}/result/{}", self.config.base_url, task_id);
168 let response = get_request_with_conn_close(
169 &self.http_client,
170 &url,
171 "Fetching task result",
172 use_new_conn,
173 )
174 .await?;
175
176 response
177 .json::<Value>()
178 .await
179 .map_err(|e| PdfConvertError::parse_error("task result response", e.to_string()))
180 };
181
182 retry_with_backoff(operation, &format!("get_task_result({task_id})")).await
183 }
184
185 fn build_form(
186 &self,
187 input: &InputDocument,
188 request: &DoclingConvertRequest,
189 ) -> Result<multipart::Form> {
190 let input_kind = input.kind()?;
191 let part = multipart::Part::stream(reqwest::Body::from(Bytes::clone(&input.bytes)))
192 .file_name(input.filename.clone())
193 .mime_str(&input.media_type)
194 .map_err(|e| {
195 PdfConvertError::api_error(None, format!("Failed to create multipart part: {e}"))
196 })?;
197
198 let mut form = multipart::Form::new()
199 .part("files", part)
200 .text("from_formats", input_kind.from_formats_value().to_string());
201
202 for format in &request.output_formats {
203 form = form.text("to_formats", format.as_api_value().to_string());
204 }
205
206 if let Some((start_page, end_page)) = request.page_range {
207 form = form.text("page_range", start_page.to_string());
208 form = form.text("page_range", end_page.to_string());
209 }
210
211 if request.chunking {
212 form = form.text("include_chunking", "true");
213 }
214
215 if matches!(
216 input_kind,
217 InputKind::Pdf | InputKind::Docx | InputKind::Markdown
218 ) {
219 form = self.apply_vlm_config(form)?;
220 }
221
222 Ok(form)
223 }
224
225 fn apply_vlm_config(&self, mut form: multipart::Form) -> Result<multipart::Form> {
226 let api_key = self.config.api_key.as_ref().ok_or_else(|| {
227 PdfConvertError::env_error(
228 "OPENAI_API_KEY",
229 "Required for Docling conversions. Please set this environment variable.",
230 )
231 })?;
232
233 let picture_description_custom_config =
234 PictureDescriptionVlmEngineOptions::for_openai_compatible(
235 &self.config.openai_base_url,
236 api_key,
237 &self.config.picture_description_model,
238 "Describe this image in a few sentences.",
239 300,
240 60,
241 );
242 let code_formula_custom_config = CodeFormulaVlmOptions {
243 scale: Some(2.0),
244 max_size: None,
245 extract_code: Some(true),
246 extract_formulas: Some(true),
247 engine_options: OpenRouterConfigBuilder::engine_options(
248 &self.config.openai_base_url,
249 api_key,
250 &self.config.code_formula_model,
251 30,
252 2,
253 ),
254 model_spec: OpenRouterConfigBuilder::model_spec(
255 &self.config.code_formula_model,
256 "Recognize code blocks and mathematical formulas in the image. For code, output the full code; for mathematical formulas, output in LaTeX format.",
257 1000,
258 ),
259 };
260 let vlm_pipeline_custom_config = VlmConvertOptions {
261 engine_options: OpenRouterConfigBuilder::engine_options(
262 &self.config.openai_base_url,
263 api_key,
264 &self.config.vlm_pipeline_model,
265 30,
266 2,
267 ),
268 model_spec: OpenRouterConfigBuilder::model_spec(
269 &self.config.vlm_pipeline_model,
270 "",
271 1000,
272 ),
273 scale: Some(1.0),
274 max_size: None,
275 batch_size: None,
276 force_backend_text: true,
277 };
278
279 form = form.text(
280 "vlm_pipeline_custom_config",
281 serde_json::to_string(&vlm_pipeline_custom_config)?,
282 );
283 form = form.text(
284 "picture_description_custom_config",
285 serde_json::to_string(&picture_description_custom_config)?,
286 );
287 form = form.text(
288 "code_formula_custom_config",
289 serde_json::to_string(&code_formula_custom_config)?,
290 );
291 form = form.text("do_code_enrichment", "True");
292 form = form.text("do_formula_enrichment", "True");
293 form = form.text("do_picture_description", "True");
294 form = form.text("ocr_engine", "rapidocr");
295 form = form.text("image_export_mode", "placeholder");
296
297 Ok(form)
298 }
299}
300
301#[cfg(test)]
302mod tests {
303 use super::*;
304
305 #[test]
306 fn build_form_uses_input_media_type_and_format() {
307 let client = DoclingClient::new(DoclingConfig {
308 base_url: "http://localhost:5001/v1".to_string(),
309 openai_base_url: "http://localhost:1234/v1".to_string(),
310 vlm_pipeline_model: "vlm".to_string(),
311 picture_description_model: "pic".to_string(),
312 code_formula_model: "code".to_string(),
313 api_key: Some("secret".to_string()),
314 })
315 .unwrap();
316
317 let input = InputDocument::new("notes.md", "text/markdown", Bytes::from_static(b"# hello"));
318 let request = DoclingConvertRequest {
319 output_formats: vec![OutputFormat::Md, OutputFormat::Text],
320 page_range: None,
321 chunking: false,
322 };
323
324 let form = client.build_form(&input, &request).unwrap();
325 let debug = format!("{form:?}");
326 assert!(debug.contains("text/markdown"));
327 assert!(debug.contains("notes.md"));
328 assert!(debug.contains("to_formats"));
329 }
330
331 #[test]
332 fn build_form_skips_page_range_for_generic_requests() {
333 let client = DoclingClient::new(DoclingConfig {
334 base_url: "http://localhost:5001/v1".to_string(),
335 openai_base_url: "http://localhost:1234/v1".to_string(),
336 vlm_pipeline_model: "vlm".to_string(),
337 picture_description_model: "pic".to_string(),
338 code_formula_model: "code".to_string(),
339 api_key: Some("secret".to_string()),
340 })
341 .unwrap();
342
343 let input = InputDocument::new(
344 "doc.docx",
345 "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
346 Bytes::from_static(b"PK"),
347 );
348 let request = DoclingConvertRequest::for_outputs(vec![OutputFormat::Md]);
349
350 let form = client.build_form(&input, &request).unwrap();
351 let debug = format!("{form:?}");
352 assert!(!debug.contains("page_range"));
353 assert!(debug.contains("from_formats"));
354 }
355}