Skip to main content

vtcode_llm/providers/
huggingface.rs

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        // Pinned so the stream-timeout fallback cannot silently regress.
741        true
742    }
743
744    fn supports_reasoning(&self, model: &str) -> bool {
745        // Codex-inspired robustness: Setting model_supports_reasoning to false
746        // does NOT disable it for known reasoning models.
747        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        // Same robustness logic for reasoning effort
757        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                    // Borrow the line from `buffer` instead of copying it into a
903                    // `String` per SSE event. The `drain` is deferred until
904                    // `serde_json::from_str` produces an owned `Value`, so the
905                    // borrow (`line`/`data`) is dead before `buffer` mutates.
906                    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                    // `event` is now owned; all borrows of `buffer` (via
935                    // `line`/`data`) are dead. Safe to drain the consumed line.
936                    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}