Skip to main content

vtcode_llm/providers/
nvidia.rs

1//! NVIDIA NIM OpenAI-compatible provider.
2
3use serde_json::{Map, Value};
4use vtcode_config::constants::{env_vars, models, urls};
5use vtcode_config::types::ReasoningEffortLevel;
6
7use super::extract_reasoning_trace;
8use super::openai_compat::{OpenAiCompatCore, OpenAiCompatSpec, impl_openai_compat_provider};
9use crate::provider::{LLMError, LLMRequest};
10
11pub struct NvidiaSpec;
12
13fn nvidia_reasoning(message: &Value, choice: &Value) -> Option<String> {
14    message
15        .get("reasoning_content")
16        .and_then(extract_reasoning_trace)
17        .or_else(|| choice.get("reasoning_content").and_then(extract_reasoning_trace))
18}
19
20impl OpenAiCompatSpec for NvidiaSpec {
21    const NAME: &'static str = "NVIDIA";
22    const KEY: &'static str = "nvidia";
23    const API_KEY_ENV: &'static str = "NVIDIA_API_KEY";
24    const DEFAULT_MODEL: &'static str = models::nvidia::DEFAULT_MODEL;
25    const DEFAULT_BASE_URL: &'static str = urls::NVIDIA_API_BASE;
26    const BASE_URL_ENV: Option<&'static str> = Some(env_vars::NVIDIA_BASE_URL);
27    const LISTED_MODELS: &'static [&'static str] = models::nvidia::SUPPORTED_MODELS;
28    // NVIDIA exposes a larger catalog than the curated VT Code picker. An
29    // explicit model selection must pass through without local rejection.
30    const VALIDATION_ALLOWLIST: Option<&'static [&'static str]> = None;
31    const STREAM_OPTIONS_INCLUDE_USAGE: bool = true;
32    const RESPONSE_REASONING_EXTRACTOR: Option<super::openai_compat::ReasoningExtractor> = Some(nvidia_reasoning);
33    const SUPPRESS_SAMPLING_WHEN_REASONING: bool = false;
34
35    fn resolve_api_key(api_key: Option<String>) -> String {
36        api_key
37            .or_else(|| std::env::var(Self::API_KEY_ENV).ok().filter(|key| !key.trim().is_empty()))
38            .unwrap_or_default()
39    }
40
41    fn insert_reasoning(
42        _core: &OpenAiCompatCore<Self>,
43        request: &LLMRequest,
44        payload: &mut Map<String, Value>,
45    ) -> Result<(), LLMError> {
46        let enable_thinking = request
47            .reasoning_effort
48            .is_some_and(|effort| effort != ReasoningEffortLevel::None);
49        payload.insert("chat_template_kwargs".to_owned(), serde_json::json!({"enable_thinking": enable_thinking}));
50        Ok(())
51    }
52
53    fn finish_payload(
54        _core: &OpenAiCompatCore<Self>,
55        request: &LLMRequest,
56        payload: &mut Map<String, Value>,
57    ) -> Result<(), LLMError> {
58        if request.tools.as_ref().is_some_and(|tools| !tools.is_empty())
59            && let Some(kwargs) = payload.get_mut("chat_template_kwargs").and_then(Value::as_object_mut)
60        {
61            kwargs.insert("force_nonempty_content".to_owned(), Value::Bool(true));
62        }
63        Ok(())
64    }
65}
66
67impl_openai_compat_provider!(NvidiaProvider, NvidiaSpec, {
68    fn supports_streaming(&self) -> bool {
69        true
70    }
71
72    fn supports_structured_output(&self, _model: &str) -> bool {
73        true
74    }
75
76    fn supports_reasoning(&self, model: &str) -> bool {
77        self.core
78            .model_behavior
79            .as_ref()
80            .and_then(|behavior| behavior.model_supports_reasoning)
81            .unwrap_or_else(|| models::nvidia::REASONING_MODELS.contains(&model) || !model.trim().is_empty())
82    }
83
84    fn supports_reasoning_effort(&self, _model: &str) -> bool {
85        true
86    }
87
88    fn effective_context_size(&self, _model: &str) -> usize {
89        1_000_000
90    }
91});
92
93#[cfg(test)]
94mod tests {
95    use super::{NvidiaProvider, NvidiaSpec};
96    use crate::BackendKind;
97    use crate::provider::{LLMProvider, LLMRequest, LLMStreamEvent, Message, ToolDefinition};
98    use crate::providers::common::parse_response_openai_format;
99    use crate::providers::openai_compat::OpenAiCompatSpec;
100    use crate::providers::shared::{OpenAiDeltaOrder, StreamAggregator, handle_openai_compatible_chunk};
101    use serde_json::json;
102    use std::sync::Arc;
103    use vtcode_config::constants::{models, urls};
104    use vtcode_config::types::ReasoningEffortLevel;
105
106    fn provider() -> NvidiaProvider {
107        NvidiaProvider::from_config(
108            Some("test-key".to_string()),
109            Some(models::nvidia::DEFAULT_MODEL.to_string()),
110            None,
111            None,
112            None,
113            None,
114            None,
115        )
116    }
117
118    fn base_request() -> LLMRequest {
119        LLMRequest {
120            messages: vec![Message::user("hello".to_string())].into(),
121            model: models::nvidia::DEFAULT_MODEL.to_string(),
122            max_tokens: Some(512),
123            temperature: Some(1.0),
124            top_p: Some(0.95),
125            stream: true,
126            ..Default::default()
127        }
128    }
129
130    #[test]
131    fn default_config_uses_nvidia_endpoint_and_bearer_key_identity() {
132        let provider = provider();
133        assert_eq!(provider.core.base_url, urls::NVIDIA_API_BASE);
134        assert_eq!(provider.core.api_key, "test-key");
135        assert_eq!(NvidiaSpec::API_KEY_ENV, "NVIDIA_API_KEY");
136        assert_eq!(provider.backend_kind(), BackendKind::Nvidia);
137
138        let overridden = NvidiaProvider::from_config(
139            Some("test-key".to_string()),
140            Some(models::nvidia::DEFAULT_MODEL.to_string()),
141            Some("https://nvidia-proxy.example/v1".to_string()),
142            None,
143            None,
144            None,
145            None,
146        );
147        assert_eq!(overridden.core.base_url, "https://nvidia-proxy.example/v1");
148    }
149
150    #[test]
151    fn golden_payload_includes_stream_usage_and_thinking_disabled_by_default() {
152        let payload = provider()
153            .core
154            .convert_request(&base_request())
155            .expect("payload should be valid");
156
157        assert_eq!(payload["model"], models::nvidia::DEFAULT_MODEL);
158        assert_eq!(payload["stream"], true);
159        assert_eq!(payload["stream_options"]["include_usage"], true);
160        assert_eq!(payload["chat_template_kwargs"]["enable_thinking"], false);
161        assert_eq!(payload["temperature"], 1.0);
162        let top_p = payload["top_p"].as_f64().expect("top_p should be numeric");
163        assert!((top_p - 0.95).abs() < 1e-6);
164    }
165
166    #[test]
167    fn reasoning_effort_toggles_nvidia_thinking() {
168        let provider = provider();
169
170        let mut request = base_request();
171        request.reasoning_effort = Some(ReasoningEffortLevel::Low);
172        let payload = provider.core.convert_request(&request).expect("payload should be valid");
173        assert_eq!(payload["chat_template_kwargs"]["enable_thinking"], true);
174
175        request.reasoning_effort = Some(ReasoningEffortLevel::None);
176        let payload = provider.core.convert_request(&request).expect("payload should be valid");
177        assert_eq!(payload["chat_template_kwargs"]["enable_thinking"], false);
178    }
179
180    #[test]
181    fn tools_force_nonempty_content_in_chat_template_kwargs() {
182        let provider = provider();
183        let mut request = base_request();
184        request.tools = Some(Arc::new(vec![ToolDefinition::function(
185            "get_weather".to_string(),
186            "Get weather".to_string(),
187            json!({"type": "object", "properties": {"city": {"type": "string"}}}),
188        )]));
189
190        let payload = provider.core.convert_request(&request).expect("payload should be valid");
191        assert_eq!(payload["chat_template_kwargs"]["force_nonempty_content"], true);
192        assert_eq!(payload["tools"][0]["type"], "function");
193    }
194
195    #[test]
196    fn arbitrary_explicit_nvidia_models_are_not_rejected() {
197        let provider = provider();
198        let request = LLMRequest {
199            model: "nvidia/custom-agent-model".to_string(),
200            messages: vec![Message::user("hello".to_string())].into(),
201            ..Default::default()
202        };
203
204        provider
205            .validate_request(&request)
206            .expect("NVIDIA should accept explicit catalog models");
207    }
208
209    #[test]
210    fn non_streaming_reasoning_content_is_extracted() {
211        let response = parse_response_openai_format::<fn(&serde_json::Value, &serde_json::Value) -> Option<String>>(
212            json!({
213                "choices": [{
214                    "message": {
215                        "content": "answer",
216                        "reasoning_content": "think first"
217                    },
218                    "finish_reason": "stop"
219                }]
220            }),
221            "NVIDIA",
222            models::nvidia::DEFAULT_MODEL.to_string(),
223            false,
224            Some(super::nvidia_reasoning),
225        )
226        .expect("response should parse");
227
228        assert_eq!(response.content.as_deref(), Some("answer"));
229        assert_eq!(response.reasoning.as_deref(), Some("think first"));
230    }
231
232    #[test]
233    fn streaming_reasoning_content_is_extracted() {
234        let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
235        let mut aggregator = StreamAggregator::new(models::nvidia::DEFAULT_MODEL.to_string());
236        let chunk = json!({"choices": [{"delta": {"reasoning_content": "think"}}]});
237
238        handle_openai_compatible_chunk(
239            &chunk,
240            &mut aggregator,
241            &tx,
242            NvidiaSpec::STREAM_REASONING_FIELDS,
243            OpenAiDeltaOrder::ReasoningFirst,
244            false,
245        );
246
247        match rx
248            .try_recv()
249            .expect("reasoning event expected")
250            .expect("stream event should be valid")
251        {
252            LLMStreamEvent::Reasoning { delta } => assert_eq!(delta, "think"),
253            other => panic!("expected reasoning event, got {other:?}"),
254        }
255    }
256}