1use std::time::{Duration, Instant};
2
3use async_openai::types::chat::{
4 ChatCompletionNamedToolChoice, ChatCompletionRequestSystemMessageArgs, ChatCompletionRequestUserMessageArgs, ChatCompletionToolChoiceOption, ChatCompletionTools, CreateChatCompletionRequestArgs
5};
6use async_openai::config::OpenAIConfig;
7use async_openai::Client;
8use async_openai::error::OpenAIError;
9use anyhow::{anyhow, Context, Result};
10use reqwest;
11use futures::future::join_all;
12
13use crate::{commit, config, debug_output, function_calling, profile};
14use crate::model::Model;
15use crate::config::AppConfig;
16use crate::multi_step_integration::generate_commit_message_multi_step;
17
18const MAX_ATTEMPTS: usize = 3;
19
20#[derive(Debug, Clone, PartialEq)]
21pub struct Response {
22 pub response: String
23}
24
25#[derive(Debug, Clone, PartialEq)]
26pub struct Request {
27 pub prompt: String,
28 pub system: String,
29 pub max_tokens: u16,
30 pub model: Model
31}
32
33pub async fn generate_commit_message(diff: &str) -> Result<String> {
36 profile!("Generate commit message (simplified)");
37
38 if let Ok(api_key) = std::env::var("OPENAI_API_KEY") {
40 if !api_key.is_empty() {
41 match commit::generate(diff.to_string(), 256, Model::GPT41Mini, None).await {
43 Ok(response) => return Ok(response.response.trim().to_string()),
44 Err(e) => {
45 log::warn!("Direct generation failed, falling back to local: {e}");
46 }
47 }
48 }
49 }
50
51 let mut lines_added = 0;
54 let mut lines_removed = 0;
55 let mut files_mentioned = std::collections::HashSet::new();
56
57 for line in diff.lines() {
58 if line.starts_with("diff --git") {
59 let parts: Vec<&str> = line.split_whitespace().collect();
61 if parts.len() >= 4 {
62 let path = parts[3].trim_start_matches("b/");
63 files_mentioned.insert(path);
64 }
65 } else if line.starts_with("+++") || line.starts_with("---") {
66 if let Some(file) = line.split_whitespace().nth(1) {
67 let cleaned = file.trim_start_matches("a/").trim_start_matches("b/");
68 if cleaned != "/dev/null" {
69 files_mentioned.insert(cleaned);
70 }
71 }
72 } else if line.starts_with('+') && !line.starts_with("+++") {
73 lines_added += 1;
74 } else if line.starts_with('-') && !line.starts_with("---") {
75 lines_removed += 1;
76 }
77 }
78
79 if let Some(session) = debug_output::debug_session() {
81 session.set_total_files_parsed(files_mentioned.len());
82 }
83
84 let message = match files_mentioned.len().cmp(&1) {
86 std::cmp::Ordering::Equal => {
87 let file = files_mentioned
88 .iter()
89 .next()
90 .ok_or_else(|| anyhow::anyhow!("No files mentioned in commit message"))?;
91 if lines_added > 0 && lines_removed == 0 {
92 format!(
93 "Add {} to {}",
94 if lines_added == 1 {
95 "content"
96 } else {
97 "new content"
98 },
99 file
100 )
101 } else if lines_removed > 0 && lines_added == 0 {
102 format!("Remove content from {file}")
103 } else {
104 format!("Update {file}")
105 }
106 }
107 std::cmp::Ordering::Greater => format!("Update {} files", files_mentioned.len()),
108 std::cmp::Ordering::Less => "Update files".to_string()
109 };
110
111 Ok(message.trim().to_string())
112}
113
114pub fn create_openai_config(settings: &AppConfig) -> Result<OpenAIConfig> {
116 let base_url = settings
118 .openai_base_url
119 .as_deref()
120 .map(str::trim)
121 .filter(|s| !s.is_empty());
122
123 let api_key = settings.openai_api_key.as_deref().unwrap_or("").trim();
124 let key_missing = api_key.is_empty() || api_key == "<PLACE HOLDER FOR YOUR API KEY>";
125
126 let effective_key = if key_missing {
130 match base_url {
131 Some(_) => "sk-no-key-required",
132 None => return Err(anyhow!("OpenAI API key not configured"))
133 }
134 } else {
135 api_key
136 };
137
138 let mut config = OpenAIConfig::new().with_api_key(effective_key);
139 if let Some(base_url) = base_url {
140 config = config.with_api_base(base_url);
141 }
142
143 Ok(config)
144}
145
146#[derive(Debug, PartialEq, Eq)]
148pub enum ModelVerification {
149 Acceptable,
151 Absent
153}
154
155pub fn classify_model(candidate: &str, available_ids: &[String], known_or_deprecated: bool) -> ModelVerification {
163 if known_or_deprecated {
164 return ModelVerification::Acceptable;
165 }
166
167 let candidate = candidate.trim();
168 if available_ids.iter().any(|id| id == candidate) {
169 ModelVerification::Acceptable
170 } else {
171 ModelVerification::Absent
172 }
173}
174
175pub async fn verify_model_exists(settings: &AppConfig, candidate: &str, known_or_deprecated: bool) -> Result<()> {
183 if known_or_deprecated {
185 return Ok(());
186 }
187
188 let config = match create_openai_config(settings) {
190 Ok(config) => config,
191 Err(e) => {
192 log::warn!("Could not verify model '{candidate}' (no usable OpenAI config: {e}); allowing it.");
193 return Ok(());
194 }
195 };
196
197 let client = Client::with_config(config);
198 match client.models().list().await {
199 Ok(list) => {
200 let ids: Vec<String> = list.data.into_iter().map(|m| m.id).collect();
201 match classify_model(candidate, &ids, known_or_deprecated) {
202 ModelVerification::Acceptable => Ok(()),
203 ModelVerification::Absent =>
204 Err(anyhow!(
205 "Model '{candidate}' is not available at the configured endpoint. \
206 Run `git ai config set model <name>` with a model the endpoint offers."
207 )),
208 }
209 }
210 Err(e) => {
211 log::warn!("Could not verify model '{candidate}' (endpoint unreachable/unauthorized: {e}); allowing it.");
212 Ok(())
213 }
214 }
215}
216
217fn truncate_to_fit(text: &str, max_tokens: usize, model: &Model) -> Result<String> {
219 profile!("Truncate to fit");
220
221 if text.len() < 1000 {
223 return Ok(text.to_string());
224 }
225
226 let token_count = model.count_tokens(text)?;
227 if token_count <= max_tokens {
228 return Ok(text.to_string());
229 }
230
231 let char_indices: Vec<(usize, char)> = text.char_indices().collect();
233 if char_indices.is_empty() {
234 return Ok(String::new());
235 }
236
237 let mut low = 0;
239 let mut high = char_indices.len();
240 let mut best_fit = String::new();
241
242 while low < high {
243 let mid = (low + high) / 2;
244
245 let byte_index = if mid < char_indices.len() {
247 char_indices[mid].0
248 } else {
249 text.len()
250 };
251
252 let truncated = &text[..byte_index];
253
254 if let Some(last_newline_pos) = truncated.rfind('\n') {
256 let candidate = &text[..last_newline_pos];
258 let candidate_tokens = model.count_tokens(candidate)?;
259
260 if candidate_tokens <= max_tokens {
261 best_fit = candidate.to_string();
262 let next_char_idx = char_indices
264 .iter()
265 .position(|(idx, _)| *idx > last_newline_pos)
266 .unwrap_or(char_indices.len());
267 low = next_char_idx;
268 } else {
269 let newline_char_idx = char_indices
271 .iter()
272 .rposition(|(idx, _)| *idx <= last_newline_pos)
273 .unwrap_or(0);
274 high = newline_char_idx;
275 }
276 } else {
277 high = mid;
278 }
279 }
280
281 if best_fit.is_empty() {
282 model.truncate(text, max_tokens)
284 } else {
285 Ok(best_fit)
286 }
287}
288
289pub async fn call_with_config(request: Request, config: OpenAIConfig) -> Result<Response> {
291 profile!("OpenAI API call with custom config");
292
293 let client = Client::with_config(config.clone());
295 let model = request.model.to_string();
296
297 match generate_commit_message_multi_step(&client, &model, &request.prompt, config::APP_CONFIG.max_commit_length).await {
298 Ok(message) => return Ok(Response { response: message }),
299 Err(e) => {
300 if e.to_string().contains("invalid_api_key") || e.to_string().contains("Incorrect API key") {
302 return Err(e);
303 }
304 log::warn!("Multi-step approach failed, falling back to single-step: {e}");
305 }
306 }
307
308 let client = if let Some(timeout) = config::APP_CONFIG.timeout {
311 let http_client = reqwest::ClientBuilder::new()
312 .timeout(Duration::from_secs(timeout as u64))
313 .build()?;
314 Client::with_config(config).with_http_client(http_client)
315 } else {
316 Client::with_config(config)
317 };
318
319 let system_tokens = request.model.count_tokens(&request.system)?;
321 let model_context_size = request.model.context_size();
322 let available_tokens = model_context_size.saturating_sub(system_tokens + request.max_tokens as usize);
323
324 let truncated_prompt = truncate_to_fit(&request.prompt, available_tokens, &request.model)?;
326
327 let commit_tool = function_calling::create_commit_function_tool(config::APP_CONFIG.max_commit_length)?;
329
330 let chat_request = CreateChatCompletionRequestArgs::default()
331 .max_completion_tokens(request.max_tokens as u32)
332 .model(request.model.to_string())
333 .messages([
334 ChatCompletionRequestSystemMessageArgs::default()
335 .content(request.system)
336 .build()?
337 .into(),
338 ChatCompletionRequestUserMessageArgs::default()
339 .content(truncated_prompt)
340 .build()?
341 .into()
342 ])
343 .tools(vec![ChatCompletionTools::Function(commit_tool)])
344 .tool_choice(ChatCompletionToolChoiceOption::Function(ChatCompletionNamedToolChoice::from("commit")))
345 .build()?;
346
347 let mut last_error = None;
348
349 for attempt in 1..=MAX_ATTEMPTS {
350 log::debug!("OpenAI API attempt {attempt} of {MAX_ATTEMPTS}");
351
352 let api_start = Instant::now();
354
355 match client.chat().create(chat_request.clone()).await {
356 Ok(response) => {
357 let api_duration = api_start.elapsed();
358
359 if let Some(session) = debug_output::debug_session() {
361 session.set_api_duration(api_duration);
362 }
363
364 log::debug!("OpenAI API call successful on attempt {attempt}");
365
366 let choice = response
368 .choices
369 .into_iter()
370 .next()
371 .context("No response choices available")?;
372
373 if let Some(tool_calls) = &choice.message.tool_calls {
375 let tool_futures: Vec<_> = tool_calls
377 .iter()
378 .filter_map(|tool_call| {
379 match tool_call {
380 async_openai::types::chat::ChatCompletionMessageToolCalls::Function(call) if call.function.name == "commit" =>
381 Some(call.function.arguments.clone()),
382 _ => None
383 }
384 })
385 .map(|args| async move { function_calling::parse_commit_function_response(&args) })
386 .collect();
387
388 let results = join_all(tool_futures).await;
390
391 let mut commit_messages = Vec::new();
393 for (i, result) in results.into_iter().enumerate() {
394 match result {
395 Ok(commit_args) => {
396 if let Some(session) = debug_output::debug_session() {
398 session.set_commit_result(commit_args.message.clone(), commit_args.reasoning.clone());
399 session.set_files_analyzed(commit_args.clone());
400 }
401 commit_messages.push(commit_args.message);
402 }
403 Err(e) => {
404 log::warn!("Failed to parse tool call {i}: {e}");
405 }
406 }
407 }
408
409 if !commit_messages.is_empty() {
411 return Ok(Response {
413 response: commit_messages
414 .into_iter()
415 .next()
416 .ok_or_else(|| anyhow::anyhow!("No commit messages generated"))?
417 });
418 }
419 }
420
421 let content = choice
423 .message
424 .content
425 .clone()
426 .context("No response content available")?;
427
428 return Ok(Response { response: content });
429 }
430 Err(e) => {
431 last_error = Some(e);
432 log::warn!("OpenAI API attempt {attempt} failed");
433
434 if let Some(OpenAIError::ApiError(ref api_err)) = last_error.as_ref() {
436 if api_err.api_error.code.as_deref() == Some("invalid_api_key") {
437 let error_msg = format!("Invalid OpenAI API key: {}", api_err.api_error.message);
438 log::error!("{error_msg}");
439 return Err(anyhow!(error_msg));
440 }
441 }
442
443 if attempt < MAX_ATTEMPTS {
444 let delay = Duration::from_millis(500 * attempt as u64);
445 log::debug!("Retrying after {delay:?}");
446 tokio::time::sleep(delay).await;
447 }
448 }
449 }
450 }
451
452 match last_error {
454 Some(OpenAIError::ApiError(api_err)) => {
455 let error_msg = format!(
456 "OpenAI API error: {} (type: {:?}, code: {:?})",
457 api_err.api_error.message,
458 api_err.api_error.r#type.as_deref().unwrap_or("unknown"),
459 api_err.api_error.code.as_deref().unwrap_or("unknown")
460 );
461 log::error!("{error_msg}");
462 Err(anyhow!(error_msg))
463 }
464 Some(e) => {
465 log::error!("OpenAI request failed: {e}");
466 Err(anyhow!("OpenAI request failed: {}", e))
467 }
468 None => Err(anyhow!("OpenAI request failed after {} attempts", MAX_ATTEMPTS))
469 }
470}
471
472pub async fn call(request: Request) -> Result<Response> {
474 profile!("OpenAI API call");
475
476 let config = create_openai_config(&config::APP_CONFIG)?;
478
479 call_with_config(request, config).await
481}
482
483#[cfg(test)]
484mod tests {
485 use async_openai::config::{Config, OpenAIConfig};
486
487 use super::*;
488
489 fn settings_with(api_key: Option<&str>, base_url: Option<&str>) -> AppConfig {
490 AppConfig {
491 openai_api_key: api_key.map(|s| s.to_string()),
492 openai_base_url: base_url.map(|s| s.to_string()),
493 model: Some("gpt-4.1-mini".to_string()),
494 max_tokens: Some(1024),
495 max_commit_length: Some(72),
496 timeout: Some(30)
497 }
498 }
499
500 #[test]
502 fn test_create_openai_config_omits_base_when_absent() {
503 let settings = settings_with(Some("sk-test-key"), None);
504 let config = create_openai_config(&settings).unwrap();
505 let default_base = OpenAIConfig::new().api_base().to_string();
506 assert_eq!(config.api_base(), default_base);
507 }
508
509 #[test]
511 fn test_create_openai_config_applies_base_when_present() {
512 let settings = settings_with(Some("sk-test-key"), Some("http://localhost:11434/v1"));
513 let config = create_openai_config(&settings).unwrap();
514 assert_eq!(config.api_base(), "http://localhost:11434/v1");
515 }
516
517 #[test]
519 fn test_create_openai_config_ignores_empty_base() {
520 let settings = settings_with(Some("sk-test-key"), Some(""));
521 let config = create_openai_config(&settings).unwrap();
522 let default_base = OpenAIConfig::new().api_base().to_string();
523 assert_eq!(config.api_base(), default_base);
524 }
525
526 #[test]
528 fn test_classify_model_known_is_acceptable() {
529 assert_eq!(classify_model("gpt-4.1", &[], true), ModelVerification::Acceptable);
530 }
531
532 #[test]
534 fn test_classify_model_present_is_acceptable() {
535 let ids = vec!["llama3.1:8b".to_string(), "mistral".to_string()];
536 assert_eq!(classify_model("llama3.1:8b", &ids, false), ModelVerification::Acceptable);
537 }
538
539 #[test]
541 fn test_classify_model_absent() {
542 let ids = vec!["llama3.1:8b".to_string()];
543 assert_eq!(classify_model("nonexistent-model", &ids, false), ModelVerification::Absent);
544 }
545
546 #[tokio::test]
548 async fn test_verify_model_exists_allows_when_no_config() {
549 let settings = settings_with(None, None);
550 assert!(verify_model_exists(&settings, "some-model", false)
552 .await
553 .is_ok());
554 }
555
556 #[tokio::test]
558 async fn test_verify_model_exists_known_short_circuits() {
559 let settings = settings_with(None, None);
560 assert!(verify_model_exists(&settings, "gpt-4.1", true)
561 .await
562 .is_ok());
563 }
564}