1use 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 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}