1#![allow(
2 clippy::bind_instead_of_map,
3 clippy::collapsible_if,
4 reason = "Intentional compatibility, platform, or test-only suppression."
5)]
6
7use crate::error_display::format_llm_error;
8use crate::provider::{
9 LLMError, LLMErrorMetadata, LLMProvider, LLMRequest, LLMResponse, LLMStream, LLMStreamEvent, MessageRole,
10 ToolDefinition,
11};
12use crate::providers::shared::{
13 NoopStreamTelemetry, StreamTelemetry, Utf8StreamDecoder, function_output_value_from_message_content,
14};
15use async_stream::try_stream;
16use async_trait::async_trait;
17use futures::StreamExt;
18use reqwest::{Client as HttpClient, Response, StatusCode};
19use serde_json::{Value, json};
20use vtcode_commons::sanitizer::sanitize_provider_diagnostic;
21use vtcode_config::TimeoutsConfig;
22use vtcode_config::constants::{env_vars, models, urls};
23use vtcode_config::core::{AnthropicConfig, ModelConfig, PromptCachingConfig};
24
25use super::common::{
26 assistant_interleaved_history_text, ensure_model, impl_llm_client, is_minimax_m2_model, map_finish_reason_common,
27 normalize_reasoning_detail_objects, override_base_url, parse_response_openai_format, resolve_model,
28};
29use super::error_handling::{format_network_error, format_parse_error};
30
31const PROVIDER_NAME: &str = "HuggingFace";
32const PROVIDER_KEY: &str = "huggingface";
33const JSON_INSTRUCTION: &str = "Return JSON that matches the provided schema.";
34
35pub struct HuggingFaceProvider {
36 api_key: String,
37 http_client: HttpClient,
38 base_url: String,
39 model: String,
40 _timeouts: TimeoutsConfig,
41 model_behavior: Option<ModelConfig>,
42}
43
44impl HuggingFaceProvider {
45 pub fn new(api_key: String) -> Self {
46 Self::with_model_internal(api_key, models::huggingface::DEFAULT_MODEL.to_string(), None, None, None)
47 }
48
49 fn with_model(api_key: String, model: String) -> Self {
50 Self::with_model_internal(api_key, model, None, None, None)
51 }
52
53 pub fn with_timeouts(api_key: String, timeouts: TimeoutsConfig) -> Self {
54 Self::with_model_internal(api_key, models::huggingface::DEFAULT_MODEL.to_string(), None, Some(timeouts), None)
55 }
56
57 fn with_model_internal(
58 api_key: String,
59 model: String,
60 base_url: Option<String>,
61 timeouts: Option<TimeoutsConfig>,
62 model_behavior: Option<ModelConfig>,
63 ) -> Self {
64 use crate::http_client::HttpClientFactory;
65
66 let timeouts = timeouts.unwrap_or_default();
67
68 Self {
69 api_key,
70 http_client: HttpClientFactory::for_llm(&timeouts),
71 base_url: override_base_url(urls::HUGGINGFACE_API_BASE, base_url, Some(env_vars::HUGGINGFACE_BASE_URL)),
72 model,
73 _timeouts: timeouts,
74 model_behavior,
75 }
76 }
77
78 pub fn from_config(
79 api_key: Option<String>,
80 model: Option<String>,
81 base_url: Option<String>,
82 _prompt_cache: Option<PromptCachingConfig>,
83 timeouts: Option<TimeoutsConfig>,
84 _anthropic: Option<AnthropicConfig>,
85 model_behavior: Option<ModelConfig>,
86 ) -> Self {
87 let api_key_value = api_key.unwrap_or_default();
88 let model_value = resolve_model(model, models::huggingface::DEFAULT_MODEL);
89 Self::with_model_internal(api_key_value, model_value, base_url, timeouts, model_behavior)
90 }
91
92 fn normalize_model_id(&self, model: &str) -> Result<String, LLMError> {
93 let model = model.trim();
94 let lower = model.to_ascii_lowercase();
95
96 if lower.contains("minimax-m2") && !model.contains(':') {
97 return Err(LLMError::Provider {
98 message: format_llm_error(
99 PROVIDER_NAME,
100 "MiniMax models require explicit provider selection (:novita suffix). \n Use 'MiniMaxAI/MiniMax-M2.5:novita'.",
101 ),
102 metadata: None,
103 });
104 }
105
106 if lower.contains("glm-5") && !model.contains(':') {
107 return Err(LLMError::Provider {
108 message: format_llm_error(
109 PROVIDER_NAME,
110 "GLM models require explicit provider selection on HuggingFace.",
111 ),
112 metadata: None,
113 });
114 }
115
116 Ok(model.to_string())
117 }
118
119 fn serialize_tools_huggingface(&self, tools: &[ToolDefinition]) -> Option<Vec<Value>> {
120 crate::providers::common::serialize_tools_openai_format(tools)
121 }
122
123 fn serialize_messages_huggingface_chat(&self, request: &LLMRequest) -> Result<Vec<Value>, LLMError> {
124 use serde_json::{Map, json};
125
126 let mut messages = Vec::with_capacity(request.messages.len());
127
128 for message in request.messages.iter() {
129 message
130 .validate_for_provider(PROVIDER_KEY)
131 .map_err(|e| LLMError::InvalidRequest { message: e, metadata: None })?;
132
133 let mut message_map = Map::with_capacity(4);
134 message_map.insert("role".to_owned(), Value::String(message.role.as_generic_str().to_owned()));
135
136 if let Some(interleaved_content) = assistant_interleaved_history_text(message, &request.model) {
137 message_map.insert("content".to_owned(), Value::String(interleaved_content));
138 } else {
139 match &message.content {
140 crate::provider::MessageContent::Text(text) => {
141 message_map.insert("content".to_owned(), Value::String(text.clone()));
142 }
143 crate::provider::MessageContent::Parts(parts) => {
144 let has_images = parts.iter().any(crate::provider::ContentPart::is_image);
145 if has_images {
146 let parts_json: Vec<Value> = parts
147 .iter()
148 .map(|part| match part {
149 crate::provider::ContentPart::Text { text } => {
150 json!({ "type": "text", "text": text })
151 }
152 crate::provider::ContentPart::Image {
153 data,
154 mime_type,
155 ..
156 } => {
157 json!({
158 "type": "image_url",
159 "image_url": {
160 "url": format!("data:{};base64,{}", mime_type, data)
161 }
162 })
163 }
164 crate::provider::ContentPart::File {
165 filename,
166 file_id,
167 file_url,
168 ..
169 } => {
170 let fallback = filename
171 .clone()
172 .or_else(|| file_id.clone())
173 .or_else(|| file_url.clone())
174 .unwrap_or_else(|| "attached file".to_string());
175 json!({ "type": "text", "text": format!("[File input not directly supported: {}]", fallback) })
176 }
177 })
178 .collect();
179 message_map.insert("content".to_owned(), Value::Array(parts_json));
180 } else {
181 let text = message.content.as_text().into_owned();
182 message_map.insert("content".to_owned(), Value::String(text));
183 }
184 }
185 }
186 }
187
188 if let Some(tool_calls) = &message.tool_calls {
189 let serialized_calls = tool_calls
190 .iter()
191 .filter_map(|call| {
192 call.function.as_ref().map(|func| {
193 json!({
194 "id": &call.id,
195 "type": "function",
196 "function": {
197 "name": &func.name,
198 "arguments": &func.arguments
199 }
200 })
201 })
202 })
203 .collect::<Vec<_>>();
204 message_map.insert("tool_calls".to_owned(), Value::Array(serialized_calls));
205 }
206
207 if let Some(tool_call_id) = &message.tool_call_id {
208 message_map.insert("tool_call_id".to_owned(), Value::String(tool_call_id.clone()));
209 }
210
211 if message.role == MessageRole::Assistant
212 && is_minimax_m2_model(&request.model)
213 && let Some(reasoning_details) = &message.reasoning_details
214 && !reasoning_details.is_empty()
215 {
216 let normalized_details = normalize_reasoning_detail_objects(reasoning_details);
217 if !normalized_details.is_empty() {
218 message_map.insert("reasoning_details".to_owned(), Value::Array(normalized_details));
219 }
220 }
221
222 messages.push(Value::Object(message_map));
223 }
224
225 Ok(messages)
226 }
227
228 fn format_for_chat_completions(&self, request: &LLMRequest) -> Result<Value, LLMError> {
229 let mut messages = self.serialize_messages_huggingface_chat(request)?;
230 let is_glm = self.is_glm_model(&request.model);
231
232 if let Some(system) = &request.system_prompt {
233 let has_system = messages.first().and_then(|m| m.get("role")).and_then(|r| r.as_str()) == Some("system");
234 if !has_system {
235 messages.insert(
236 0,
237 json!({
238 "role": "system",
239 "content": system
240 }),
241 );
242 }
243 }
244
245 let mut payload = json!({
246 "model": request.model,
247 "messages": messages,
248 "stream": request.stream,
249 });
250
251 if request.stream && request.tools.is_some() && is_glm {
252 payload["tool_stream"] = json!(true);
253 }
254
255 if let Some(max_tokens) = request.max_tokens {
256 payload["max_tokens"] = json!(max_tokens);
257 }
258
259 if let Some(tools) = &request.tools {
260 if let Some(serialized) = self.serialize_tools_huggingface(tools) {
261 payload["tools"] = json!(serialized);
262
263 if let Some(choice) = &request.tool_choice {
264 payload["tool_choice"] = choice.to_provider_format("openai");
265 }
266 }
267 }
268
269 if let Some(temperature) = request.temperature {
270 payload["temperature"] = json!(super::common::sampling_param_f64(temperature));
271 }
272
273 if let Some(top_p) = request.top_p {
274 payload["top_p"] = json!(super::common::sampling_param_f64(top_p));
275 }
276
277 if let Some(top_k) = request.top_k {
278 payload["top_k"] = json!(top_k);
279 }
280
281 if let Some(effort) = request.reasoning_effort {
282 use crate::rig_adapter::RigProviderCapabilities;
283 use vtcode_config::models::Provider;
284 let supported = self.supported_reasoning_efforts(&request.model);
285 if let Some(reasoning_params) = RigProviderCapabilities::new(Provider::HuggingFace, &request.model)
286 .reasoning_parameters_for_supported_efforts(effort, supported)?
287 {
288 if let Some(params_obj) = reasoning_params.as_object() {
289 for (k, v) in params_obj {
290 payload[k] = v.clone();
291 }
292 }
293 }
294 }
295
296 if request.output_format.is_some() && !is_glm {
297 payload["response_format"] = json!({ "type": "json_object" });
298 }
299
300 Ok(payload)
301 }
302
303 fn is_glm_model(&self, model: &str) -> bool {
304 let lower = model.to_ascii_lowercase();
305 lower.contains("glm")
306 }
307
308 fn is_deepseek_model(&self, model: &str) -> bool {
309 let lower = model.to_ascii_lowercase();
310 lower.contains("deepseek")
311 }
312
313 fn is_minimax_model(&self, model: &str) -> bool {
314 let lower = model.to_ascii_lowercase();
315 lower.contains("minimax")
316 }
317
318 fn apply_model_defaults(&self, request: &mut LLMRequest) {
319 if self.is_minimax_model(&request.model) {
320 if request.temperature.is_none() {
321 request.temperature = Some(1.0);
322 }
323 if request.top_p.is_none() {
324 request.top_p = Some(0.95);
325 }
326 if request.top_k.is_none() {
327 request.top_k = Some(40);
328 }
329 }
330 }
331
332 fn add_json_instruction(&self, payload: &mut Value) -> Result<(), LLMError> {
333 if let Some(instructions) = payload.get_mut("instructions") {
334 if let Some(text) = instructions.as_str() {
335 if !text.contains("Return JSON") {
336 *instructions = json!(format!("{}\n\n{}", text, JSON_INSTRUCTION));
337 }
338 }
339 } else {
340 payload["instructions"] = json!(JSON_INSTRUCTION);
341 }
342
343 Ok(())
344 }
345
346 fn format_for_responses_api(&self, request: &LLMRequest) -> Result<Value, LLMError> {
347 let mut input = Vec::new();
348
349 for msg in request.messages.iter() {
350 let convert_parts = |parts: &[crate::provider::ContentPart]| -> Value {
351 let parts_json: Vec<Value> = parts
352 .iter()
353 .map(|part| match part {
354 crate::provider::ContentPart::Text { text } => {
355 json!({ "type": "input_text", "text": text })
356 }
357 crate::provider::ContentPart::Image { data, mime_type, .. } => {
358 json!({
359 "type": "input_image",
360 "image_url": format!("data:{};base64,{}", mime_type, data)
361 })
362 }
363 crate::provider::ContentPart::File { filename, file_id, file_url, .. } => {
364 let fallback = filename
365 .clone()
366 .or_else(|| file_id.clone())
367 .or_else(|| file_url.clone())
368 .unwrap_or_else(|| "attached file".to_string());
369 json!({
370 "type": "input_text",
371 "text": format!("[File input not directly supported: {}]", fallback)
372 })
373 }
374 })
375 .collect();
376 json!(parts_json)
377 };
378
379 match msg.role {
380 MessageRole::System | MessageRole::User => {
381 if msg.role == MessageRole::System && request.system_prompt.is_some() {
382 if let crate::provider::MessageContent::Text(text) = &msg.content {
383 if request.system_prompt.as_ref().map(|s| s.as_ref()) == Some(text.as_str()) {
384 continue;
385 }
386 }
387 }
388
389 let role = if msg.role == MessageRole::System {
390 "system"
391 } else {
392 "user"
393 };
394
395 let mut message_obj = json!({
396 "type": "message",
397 "role": role,
398 });
399
400 match &msg.content {
401 crate::provider::MessageContent::Text(text) => {
402 message_obj["content"] = json!(text);
403 }
404 crate::provider::MessageContent::Parts(parts) => {
405 message_obj["content"] = convert_parts(parts);
406 }
407 }
408
409 input.push(message_obj);
410 }
411 MessageRole::Assistant => {
412 let has_content = match &msg.content {
413 crate::provider::MessageContent::Text(text) => !text.is_empty(),
414 crate::provider::MessageContent::Parts(parts) => !parts.is_empty(),
415 };
416
417 if has_content {
418 let mut message_obj = json!({
419 "type": "message",
420 "role": "assistant",
421 });
422
423 match &msg.content {
424 crate::provider::MessageContent::Text(text) => {
425 message_obj["content"] = json!(text);
426 }
427 crate::provider::MessageContent::Parts(parts) => {
428 message_obj["content"] = convert_parts(parts);
429 }
430 }
431
432 input.push(message_obj);
433 }
434
435 if let Some(tool_calls) = &msg.tool_calls {
436 for tc in tool_calls {
437 if let Some(func) = &tc.function {
438 input.push(json!({
439 "type": "function_call",
440 "call_id": tc.id,
441 "name": func.name,
442 "arguments": func.arguments
443 }));
444 }
445 }
446 }
447 }
448 MessageRole::Tool => {
449 input.push(json!({
450 "type": "function_call_output",
451 "call_id": msg.tool_call_id.clone().unwrap_or_default(),
452 "output": function_output_value_from_message_content(&msg.content)
453 }));
454 }
455 }
456 }
457
458 let mut payload = json!({
459 "model": request.model,
460 "input": input,
461 "stream": request.stream,
462 });
463
464 if let Some(system_prompt) = &request.system_prompt {
465 payload["instructions"] = json!(system_prompt);
466 }
467
468 if let Some(effort) = request.reasoning_effort {
469 use vtcode_config::types::ReasoningEffortLevel;
470 if effort != ReasoningEffortLevel::None {
471 payload["reasoning"] = json!({ "effort": effort.as_str() });
472 }
473 }
474
475 if let Some(max_tokens) = request.max_tokens {
476 payload["max_tokens"] = json!(max_tokens);
477 }
478 if let Some(temperature) = request.temperature {
479 payload["temperature"] = json!(super::common::sampling_param_f64(temperature));
480 }
481 if let Some(top_p) = request.top_p {
482 payload["top_p"] = json!(super::common::sampling_param_f64(top_p));
483 }
484 if let Some(top_k) = request.top_k {
485 payload["top_k"] = json!(top_k);
486 }
487
488 if let Some(tools) = &request.tools {
489 if let Some(serialized) = self.serialize_tools_huggingface(tools) {
490 payload["tools"] = json!(serialized);
491
492 if let Some(choice) = &request.tool_choice {
493 payload["tool_choice"] = choice.to_provider_format("openai");
494 }
495 }
496 }
497
498 if request.output_format.is_some() || request.tools.is_some() {
499 self.add_json_instruction(&mut payload)?;
500 }
501
502 if request.output_format.is_some() && !self.is_glm_model(&request.model) {
503 payload["response_format"] = json!({ "type": "json_object" });
504 }
505
506 Ok(payload)
507 }
508
509 fn should_use_responses_api(&self, _request: &LLMRequest) -> bool {
510 false
511 }
512
513 fn format_error(&self, status: StatusCode, body: &str) -> LLMError {
514 let message = format!("HuggingFace API error ({status}): {body}");
515
516 LLMError::Provider {
517 message: format_llm_error(PROVIDER_NAME, &message),
518 metadata: Some(LLMErrorMetadata::new(
519 PROVIDER_NAME,
520 Some(status.as_u16()),
521 None,
522 None,
523 None,
524 None,
525 Some(sanitize_provider_diagnostic(body.as_bytes())),
526 )),
527 }
528 }
529
530 fn parse_responses_api_format(json: &Value, model: String) -> Result<LLMResponse, LLMError> {
531 let convenience_text = json.get("output_text").and_then(|t| t.as_str());
532
533 let json_obj = json.get("response").unwrap_or(json);
534
535 let output = json_obj.get("output").and_then(|v| v.as_array());
536
537 let output_arr = match output {
538 Some(arr) => arr,
539 None => {
540 if let Some(text) = convenience_text {
541 return Ok(LLMResponse {
542 content: Some(text.to_string()),
543 tool_calls: None,
544 model,
545 usage: None,
546 finish_reason: crate::provider::FinishReason::Stop,
547 reasoning: None,
548 reasoning_details: None,
549 tool_references: Vec::new(),
550 request_id: None,
551 organization_id: None,
552 compaction: None,
553 });
554 }
555
556 return Err(LLMError::Provider {
557 message: format_llm_error(PROVIDER_NAME, "Not a Responses API format"),
558 metadata: None,
559 });
560 }
561 };
562
563 let mut content_fragments: Vec<String> = Vec::new();
564 let mut reasoning_fragments: Vec<String> = Vec::new();
565 let mut tool_calls: Vec<crate::provider::ToolCall> = Vec::new();
566
567 for item in output_arr {
568 let item_type = item.get("type").and_then(|t| t.as_str()).unwrap_or("");
569
570 match item_type {
571 "message" => {
572 if let Some(content_arr) = item.get("content").and_then(|c| c.as_array()) {
573 for entry in content_arr {
574 let entry_type = entry.get("type").and_then(|t| t.as_str()).unwrap_or("");
575 match entry_type {
576 "text" | "output_text" => {
577 if let Some(text) = entry.get("text").and_then(|t| t.as_str()) {
578 if !text.is_empty() {
579 content_fragments.push(text.to_string());
580 }
581 }
582 }
583 "reasoning" => {
584 if let Some(text) = entry.get("text").and_then(|t| t.as_str()) {
585 if !text.is_empty() {
586 reasoning_fragments.push(text.to_string());
587 }
588 }
589 }
590 "function_call" | "tool_call" => {
591 if let Some(call) = Self::parse_responses_tool_call(entry) {
592 tool_calls.push(call);
593 }
594 }
595 _ => {}
596 }
597 }
598 }
599 }
600 "function_call" | "tool_call" => {
601 if let Some(call) = Self::parse_responses_tool_call(item) {
602 tool_calls.push(call);
603 }
604 }
605 "reasoning" => {
606 if let Some(summary_arr) = item.get("summary").and_then(|s| s.as_array()) {
607 for summary in summary_arr {
608 if let Some(text) = summary.get("text").and_then(|t| t.as_str()) {
609 if !text.is_empty() {
610 reasoning_fragments.push(text.to_string());
611 }
612 }
613 }
614 } else if let Some(text) = item.get("text").and_then(|t| t.as_str()) {
615 reasoning_fragments.push(text.to_string());
616 }
617 }
618 _ => {}
619 }
620 }
621
622 let content = if content_fragments.is_empty() {
623 convenience_text.map(|t| t.to_string())
624 } else {
625 Some(content_fragments.join(""))
626 };
627
628 let reasoning = if reasoning_fragments.is_empty() {
629 None
630 } else {
631 Some(reasoning_fragments.join("\n\n"))
632 };
633
634 let finish_reason = if !tool_calls.is_empty() {
635 crate::provider::FinishReason::ToolCalls
636 } else {
637 crate::provider::FinishReason::Stop
638 };
639
640 let usage_value = json.get("usage").or_else(|| json_obj.get("usage"));
641 let usage = usage_value.map(|usage_value| crate::provider::Usage {
642 prompt_tokens: usage_value
643 .get("input_tokens")
644 .or_else(|| usage_value.get("prompt_tokens"))
645 .and_then(|pt| pt.as_u64())
646 .unwrap_or(0) as u32,
647 completion_tokens: usage_value
648 .get("output_tokens")
649 .or_else(|| usage_value.get("completion_tokens"))
650 .and_then(|ct| ct.as_u64())
651 .unwrap_or(0) as u32,
652 total_tokens: usage_value.get("total_tokens").and_then(|tt| tt.as_u64()).unwrap_or(0) as u32,
653 cached_prompt_tokens: None,
654 cache_creation_tokens: None,
655 cache_read_tokens: None,
656 iterations: None,
657 });
658
659 Ok(LLMResponse {
660 content,
661 tool_calls: if tool_calls.is_empty() { None } else { Some(tool_calls) },
662 model,
663 usage,
664 finish_reason,
665 reasoning,
666 reasoning_details: None,
667 tool_references: Vec::new(),
668 request_id: None,
669 organization_id: None,
670 compaction: None,
671 })
672 }
673
674 fn parse_responses_tool_call(item: &Value) -> Option<crate::provider::ToolCall> {
675 let call_id = item.get("id").and_then(|v| v.as_str()).unwrap_or("");
676 let function_obj = item.get("function").and_then(|v| v.as_object());
677 let name = function_obj.and_then(|f| f.get("name").and_then(|n| n.as_str()))?;
678 let arguments = function_obj.and_then(|f| f.get("arguments"));
679
680 let serialized = arguments.map_or("{}".to_owned(), |args| {
681 if args.is_string() {
682 args.as_str().unwrap_or("{}").to_string()
683 } else {
684 args.to_string()
685 }
686 });
687
688 Some(crate::provider::ToolCall::function(call_id.to_string(), name.to_string(), serialized))
689 }
690
691 async fn parse_response(
692 &self,
693 response: Response,
694 model: String,
695 use_responses_api: bool,
696 ) -> Result<LLMResponse, LLMError> {
697 let status = response.status();
698
699 if !status.is_success() {
700 let body = crate::providers::common::read_provider_error_body(response).await;
701 return Err(self.format_error(status, &body));
702 }
703
704 let json: Value = response.json().await.map_err(|err| format_parse_error(PROVIDER_NAME, &err))?;
705
706 if use_responses_api {
707 if json.get("output").is_some() {
708 return Self::parse_responses_api_format(&json, model);
709 }
710 }
711
712 parse_response_openai_format::<fn(&Value, &Value) -> Option<String>>(json, PROVIDER_NAME, model, false, None)
713 }
714
715 fn available_models() -> Vec<String> {
716 models::huggingface::SUPPORTED_MODELS.iter().map(|s| s.to_string()).collect()
717 }
718
719 fn get_endpoint(&self, use_responses_api: bool) -> String {
720 let base = self.base_url.trim_end_matches('/');
721 if use_responses_api {
722 format!("{base}/responses")
723 } else {
724 super::common::chat_completions_url(base)
725 }
726 }
727}
728
729#[async_trait]
730impl LLMProvider for HuggingFaceProvider {
731 fn name(&self) -> &str {
732 PROVIDER_KEY
733 }
734
735 fn supports_streaming(&self) -> bool {
736 true
737 }
738
739 fn supports_non_streaming(&self, _model: &str) -> bool {
740 true
742 }
743
744 fn supports_reasoning(&self, model: &str) -> bool {
745 models::huggingface::REASONING_MODELS.contains(&model)
748 || self
749 .model_behavior
750 .as_ref()
751 .and_then(|b| b.model_supports_reasoning)
752 .unwrap_or(false)
753 }
754
755 fn supports_reasoning_effort(&self, model: &str) -> bool {
756 self.is_glm_model(model)
758 || self.is_deepseek_model(model)
759 || self
760 .model_behavior
761 .as_ref()
762 .and_then(|b| b.model_supports_reasoning_effort)
763 .unwrap_or(false)
764 }
765
766 fn supports_tools(&self, _model: &str) -> bool {
767 true
768 }
769
770 fn supports_parallel_tool_config(&self, _model: &str) -> bool {
771 false
772 }
773
774 fn supports_structured_output(&self, _model: &str) -> bool {
775 true
776 }
777
778 fn supports_context_caching(&self, _model: &str) -> bool {
779 false
780 }
781
782 fn effective_context_size(&self, model: &str) -> usize {
783 crate::provider::catalog_context_window("huggingface", model, 128_000)
784 }
785
786 async fn generate(&self, mut request: LLMRequest) -> Result<LLMResponse, LLMError> {
787 let model = ensure_model(&mut request, &self.model);
788
789 self.apply_model_defaults(&mut request);
790 self.validate_request(&request)?;
791
792 let model_id = self.normalize_model_id(&request.model)?;
793 request.model = model_id;
794
795 let use_responses_api = self.should_use_responses_api(&request);
796 let payload = if use_responses_api {
797 self.format_for_responses_api(&request)?
798 } else {
799 self.format_for_chat_completions(&request)?
800 };
801
802 let endpoint = self.get_endpoint(use_responses_api);
803
804 let response = self
805 .http_client
806 .post(&endpoint)
807 .header("Authorization", format!("Bearer {}", self.api_key))
808 .json(&payload)
809 .send()
810 .await
811 .map_err(|err| format_network_error(PROVIDER_NAME, &err))?;
812
813 self.parse_response(response, model, use_responses_api).await
814 }
815
816 async fn stream(&self, mut request: LLMRequest) -> Result<LLMStream, LLMError> {
817 let model = ensure_model(&mut request, &self.model);
818
819 self.apply_model_defaults(&mut request);
820 self.validate_request(&request)?;
821 request.stream = true;
822
823 let model_id = self.normalize_model_id(&request.model)?;
824 request.model = model_id;
825
826 let use_responses_api = self.should_use_responses_api(&request);
827 let payload = if use_responses_api {
828 self.format_for_responses_api(&request)?
829 } else {
830 self.format_for_chat_completions(&request)?
831 };
832
833 let endpoint = self.get_endpoint(use_responses_api);
834
835 let response = self
836 .http_client
837 .post(&endpoint)
838 .header("Authorization", format!("Bearer {}", self.api_key))
839 .json(&payload)
840 .send()
841 .await
842 .map_err(|err| format_network_error(PROVIDER_NAME, &err))?;
843
844 if !response.status().is_success() {
845 let status = response.status();
846 let body = crate::providers::common::read_provider_error_body(response).await;
847 return Err(self.format_error(status, &body));
848 }
849
850 self.create_stream(response, model, use_responses_api).await
851 }
852
853 fn supported_models(&self) -> Vec<String> {
854 Self::available_models()
855 }
856
857 fn validate_request(&self, request: &LLMRequest) -> Result<(), LLMError> {
858 if request.messages.is_empty() {
859 return Err(LLMError::InvalidRequest {
860 message: format_llm_error(PROVIDER_NAME, "Messages cannot be empty"),
861 metadata: None,
862 });
863 }
864
865 if request.model.trim().is_empty() {
866 return Err(LLMError::InvalidRequest {
867 message: format_llm_error(PROVIDER_NAME, "Model identifier cannot be empty"),
868 metadata: None,
869 });
870 }
871
872 Ok(())
873 }
874}
875
876impl HuggingFaceProvider {
877 async fn create_stream(
878 &self,
879 response: Response,
880 model: String,
881 use_responses_api: bool,
882 ) -> Result<LLMStream, LLMError> {
883 let mut bytes_stream = response.bytes_stream();
884 let mut buffer = String::with_capacity(4096);
885 let mut decoder = Utf8StreamDecoder::new();
886 let mut aggregator = crate::providers::shared::StreamAggregator::new(model.clone());
887 let telemetry = NoopStreamTelemetry;
888
889 let stream = try_stream! {
890 'outer: while let Some(chunk_result) = bytes_stream.next().await {
891 let chunk = chunk_result.map_err(|err| format_network_error(PROVIDER_NAME, &err))?;
892 buffer.push_str(&decoder.push(&chunk));
893
894 if buffer.len() > 128_000 {
895 Err(LLMError::Provider {
896 message: format_llm_error(PROVIDER_NAME, "Stream buffer exceeded maximum size (128KB)"),
897 metadata: None,
898 })?;
899 }
900
901 while let Some(newline_pos) = buffer.find('\n') {
902 let line = buffer[..newline_pos].trim();
907
908 if line.is_empty() || line.starts_with(':') {
909 buffer.drain(..=newline_pos);
910 continue;
911 }
912
913 let data = match line.strip_prefix("data: ") {
914 Some(stripped) => stripped,
915 None => {
916 buffer.drain(..=newline_pos);
917 continue;
918 }
919 };
920
921 if data == "[DONE]" {
922 buffer.drain(..=newline_pos);
923 break 'outer;
924 }
925
926 let event: Value = match serde_json::from_str(data) {
927 Ok(v) => v,
928 Err(_) => {
929 buffer.drain(..=newline_pos);
930 continue;
931 }
932 };
933
934 buffer.drain(..=newline_pos);
937
938 if use_responses_api {
939 let event_type = event.get("type").and_then(|t| t.as_str()).unwrap_or("");
940
941 match event_type {
942 "response.output_text.delta" | "output_text.delta" => {
943 if let Some(delta) = event.get("delta").and_then(|d| d.as_str()) {
944 telemetry.on_content_delta(delta);
945 for ev in aggregator.handle_content(delta) {
946 yield ev;
947 }
948 }
949 continue;
950 }
951 "response.reasoning.delta" | "reasoning.delta" => {
952 if let Some(delta) = event.get("delta").and_then(|d| d.as_str()) {
953 if let Some(d) = aggregator.handle_reasoning(delta) {
954 telemetry.on_reasoning_delta(&d);
955 yield LLMStreamEvent::Reasoning { delta: d };
956 }
957 }
958 continue;
959 }
960 "response.function_call_arguments.delta" | "tool_call.delta" => {
961 telemetry.on_tool_call_delta();
962 continue;
963 }
964 "response.completed" => {
965 if let Some(response_obj) = event.get("response") {
966 if let Ok(response) = Self::parse_responses_api_format(response_obj, model.clone()) {
967 let final_agg_response = aggregator.finalize();
968 let mut merged_response = response;
969 if merged_response.content.is_none() {
970 merged_response.content = final_agg_response.content;
971 }
972 if merged_response.reasoning.is_none() {
973 merged_response.reasoning = final_agg_response.reasoning;
974 }
975 if merged_response.tool_calls.is_none() {
976 merged_response.tool_calls = final_agg_response.tool_calls;
977 }
978 if merged_response.usage.is_none() {
979 merged_response.usage = final_agg_response.usage;
980 }
981 yield LLMStreamEvent::Completed { response: Box::new(merged_response) };
982 return;
983 }
984 }
985 break 'outer;
986 }
987 "response.done" => {
988 break 'outer;
989 }
990 _ => {}
991 }
992 }
993
994 if let Some(choices_arr) = event.get("choices").and_then(|c| c.as_array()) {
995 if let Some(choice) = choices_arr.first() {
996 if let Some(delta_obj) = choice.get("delta") {
997 if let Some(content) = delta_obj.get("content").and_then(|c| c.as_str()) {
998 telemetry.on_content_delta(content);
999 for ev in aggregator.handle_content(content) {
1000 yield ev;
1001 }
1002 }
1003
1004 if let Some(reason) = delta_obj.get("reasoning_content").and_then(|r| r.as_str()) {
1005 if let Some(d) = aggregator.handle_reasoning(reason) {
1006 telemetry.on_reasoning_delta(&d);
1007 yield LLMStreamEvent::Reasoning { delta: d };
1008 }
1009 }
1010
1011 if let Some(reasoning_details) = delta_obj
1012 .get("reasoning_details")
1013 .and_then(|details| details.as_array())
1014 {
1015 aggregator.set_reasoning_details(reasoning_details);
1016 }
1017
1018 if let Some(tool_calls_arr) = delta_obj.get("tool_calls").and_then(|tc| tc.as_array()) {
1019 aggregator.handle_tool_calls(tool_calls_arr);
1020 telemetry.on_tool_call_delta();
1021 }
1022 }
1023
1024 if let Some(finish_reason_str) = choice.get("finish_reason").and_then(|fr| fr.as_str()) {
1025 aggregator.set_finish_reason(map_finish_reason_common(finish_reason_str));
1026 if let Some(usage_value) = event.get("usage") {
1027 aggregator.set_usage(crate::provider::Usage {
1028 prompt_tokens: usage_value.get("prompt_tokens").and_then(|pt| pt.as_u64()).unwrap_or(0) as u32,
1029 completion_tokens: usage_value.get("completion_tokens").and_then(|ct| ct.as_u64()).unwrap_or(0) as u32,
1030 total_tokens: usage_value.get("total_tokens").and_then(|tt| tt.as_u64()).unwrap_or(0) as u32,
1031 cached_prompt_tokens: None,
1032 cache_creation_tokens: None,
1033 cache_read_tokens: None,
1034 iterations: None,
1035 });
1036 }
1037
1038 break 'outer;
1039 }
1040 }
1041 }
1042 }
1043 }
1044
1045 yield LLMStreamEvent::Completed { response: Box::new(aggregator.finalize()) };
1046 };
1047
1048 Ok(Box::pin(stream))
1049 }
1050}
1051
1052impl_llm_client!(HuggingFaceProvider);
1053
1054#[cfg(test)]
1055mod tests {
1056 use super::HuggingFaceProvider;
1057 use crate::provider::{LLMRequest, Message, ToolDefinition};
1058 use crate::providers::common::{is_minimax_m2_model, normalize_reasoning_detail_object};
1059 use serde_json::json;
1060 use std::sync::Arc;
1061
1062 #[test]
1063 fn minimax_model_detection_handles_variants() {
1064 assert!(is_minimax_m2_model("MiniMaxAI/MiniMax-M2.5:novita"));
1065 assert!(is_minimax_m2_model("minimax-m2.5"));
1066 assert!(!is_minimax_m2_model("deepseek-r1"));
1067 }
1068
1069 #[test]
1070 fn normalize_reasoning_detail_decodes_stringified_object() {
1071 let parsed = normalize_reasoning_detail_object(&json!("{\"type\":\"reasoning.text\",\"text\":\"step\"}"))
1072 .expect("expected a parsed reasoning detail object");
1073 assert!(parsed.is_object());
1074 assert_eq!(parsed["type"], "reasoning.text");
1075 }
1076
1077 #[test]
1078 fn serialize_messages_normalizes_minimax_reasoning_details() {
1079 let provider =
1080 HuggingFaceProvider::with_model("test-key".to_string(), "MiniMaxAI/MiniMax-M2.5:novita".to_string());
1081 let request = LLMRequest {
1082 model: "MiniMaxAI/MiniMax-M2.5:novita".to_string(),
1083 messages: vec![
1084 Message::assistant("answer".to_string())
1085 .with_reasoning_details(Some(vec![json!("{\"type\":\"reasoning.text\",\"text\":\"chain\"}")])),
1086 ]
1087 .into(),
1088 ..Default::default()
1089 };
1090
1091 let messages = provider
1092 .serialize_messages_huggingface_chat(&request)
1093 .expect("message serialization should succeed");
1094 assert!(messages[0]["reasoning_details"].is_array());
1095 assert!(messages[0]["reasoning_details"][0].is_object());
1096 }
1097
1098 #[test]
1099 fn serialize_messages_rehydrates_glm_interleaved_history_into_content() {
1100 let provider = HuggingFaceProvider::with_model("test-key".to_string(), "zai-org/GLM-5.1:novita".into());
1101 let request = LLMRequest {
1102 model: "zai-org/GLM-5.1:novita".to_string(),
1103 messages: vec![Message::assistant("done".to_string()).with_reasoning(Some("trace".to_string()))].into(),
1104 ..Default::default()
1105 };
1106
1107 let messages = provider
1108 .serialize_messages_huggingface_chat(&request)
1109 .expect("message serialization should succeed");
1110
1111 assert_eq!(messages[0]["content"], json!("<think>trace</think>done"));
1112 }
1113
1114 #[test]
1115 fn format_for_chat_completions_keeps_apply_patch_as_function_tool() {
1116 let provider =
1117 HuggingFaceProvider::with_model("test-key".to_string(), "Qwen/Qwen3-Coder-480B-A35B-Instruct".to_string());
1118 let request = LLMRequest {
1119 model: "Qwen/Qwen3-Coder-480B-A35B-Instruct".to_string(),
1120 messages: vec![Message::user("apply a patch".to_string())].into(),
1121 tools: Some(Arc::new(vec![ToolDefinition::apply_patch("Apply patches".to_string())])),
1122 ..Default::default()
1123 };
1124
1125 let payload = provider
1126 .format_for_chat_completions(&request)
1127 .expect("payload should serialize");
1128
1129 assert_eq!(payload["tools"][0]["type"], "function");
1130 assert_eq!(payload["tools"][0]["function"]["name"], "apply_patch");
1131 }
1132}