Skip to main content

vtcode_llm/providers/
copilot.rs

1use std::fmt::Write;
2use std::path::PathBuf;
3use std::sync::Arc;
4
5use anyhow::Context;
6use async_stream::stream;
7use async_trait::async_trait;
8use tokio::sync::Mutex;
9use vtcode_config::auth::CopilotAuthConfig;
10use vtcode_config::constants::models::copilot as copilot_models;
11use vtcode_config::models::supported_models_for_provider;
12
13use crate::copilot::{
14    COPILOT_MODEL_ID, COPILOT_PROVIDER_KEY, CopilotAcpClient, CopilotPromptSessionFuture, CopilotRuntimeRequest,
15    CopilotToolCallFailure, CopilotToolCallResponse, PromptSession, PromptSessionCancelHandle, PromptUpdate,
16    probe_auth_status,
17};
18use crate::provider::{
19    LLMError, LLMProvider, LLMRequest, LLMResponse, LLMStream, LLMStreamEvent, Message, MessageRole, ToolDefinition,
20};
21use crate::providers::common::validate_request_common;
22
23pub struct CopilotProvider {
24    model: String,
25    auth_config: CopilotAuthConfig,
26    workspace_root: PathBuf,
27    client: Mutex<Option<CachedCopilotClient>>,
28}
29
30struct CachedCopilotClient {
31    raw_model: Option<String>,
32    tool_signature: String,
33    client: Arc<CopilotAcpClient>,
34}
35
36#[derive(Debug, Clone, PartialEq, Eq)]
37struct ResolvedCopilotModel {
38    request_model: String,
39    raw_model: Option<String>,
40}
41
42impl CopilotProvider {
43    pub fn from_config(
44        model: Option<String>,
45        auth_config: Option<CopilotAuthConfig>,
46        workspace_root: Option<PathBuf>,
47    ) -> Self {
48        Self {
49            model: model.unwrap_or_else(|| COPILOT_MODEL_ID.to_string()),
50            auth_config: auth_config.unwrap_or_default(),
51            workspace_root: workspace_root
52                .unwrap_or_else(|| std::env::current_dir().unwrap_or_else(|_| PathBuf::from("."))),
53            client: Mutex::new(None),
54        }
55    }
56
57    async fn client(
58        &self,
59        model: &ResolvedCopilotModel,
60        tools: &[ToolDefinition],
61    ) -> Result<Arc<CopilotAcpClient>, LLMError> {
62        let tool_signature = copilot_tool_signature(tools);
63        if let Some(client) = self.cached_client(model, &tool_signature).await {
64            return Ok(client);
65        }
66
67        let auth_status = probe_auth_status(&self.auth_config, Some(&self.workspace_root)).await;
68        if !auth_status.is_authenticated() {
69            return Err(LLMError::Authentication {
70                message: auth_status
71                    .message
72                    .unwrap_or_else(|| "GitHub Copilot is not authenticated. Run `vtcode login copilot`.".to_string()),
73                metadata: None,
74            });
75        }
76
77        let llm_tools: Vec<ToolDefinition> = tools.to_vec();
78
79        let created = Arc::new(
80            CopilotAcpClient::connect(&self.auth_config, &self.workspace_root, model.raw_model.as_deref(), &llm_tools)
81                .await
82                .map_err(map_copilot_error)?,
83        );
84
85        let mut client = self.client.lock().await;
86        if let Some(existing) = client.as_ref()
87            && existing.raw_model.as_deref() == model.raw_model.as_deref()
88            && existing.tool_signature == tool_signature
89        {
90            return Ok(existing.client.clone());
91        }
92        *client = Some(CachedCopilotClient {
93            raw_model: model.raw_model.clone(),
94            tool_signature,
95            client: created.clone(),
96        });
97        Ok(created)
98    }
99
100    async fn cached_client(&self, model: &ResolvedCopilotModel, tool_signature: &str) -> Option<Arc<CopilotAcpClient>> {
101        let client = self.client.lock().await;
102        client
103            .as_ref()
104            .filter(|cached| {
105                cached.raw_model.as_deref() == model.raw_model.as_deref() && cached.tool_signature == tool_signature
106            })
107            .map(|cached| cached.client.clone())
108    }
109
110    fn resolve_model(&self, request: &LLMRequest) -> Result<ResolvedCopilotModel, LLMError> {
111        let requested = if request.model.trim().is_empty() {
112            self.model.trim()
113        } else {
114            request.model.trim()
115        };
116
117        let raw_model = normalize_copilot_model_id(requested).ok_or_else(|| {
118            invalid_request(&format!(
119                "Unsupported GitHub Copilot model: {requested}. Choose `copilot-auto` or a live GitHub Copilot model id from the picker."
120            ))
121        })?;
122
123        Ok(ResolvedCopilotModel { request_model: requested.to_string(), raw_model })
124    }
125
126    fn build_transcript(&self, request: &LLMRequest) -> Result<String, LLMError> {
127        let mut transcript = String::new();
128
129        if let Some(system_prompt) = request.system_prompt.as_ref() {
130            append_block(&mut transcript, "System", system_prompt);
131        }
132
133        for message in request.messages.iter() {
134            let label = match message.role {
135                MessageRole::System => "System",
136                MessageRole::User => "User",
137                MessageRole::Assistant => "Assistant",
138                MessageRole::Tool => "Tool",
139            };
140            append_block(&mut transcript, label, &render_message_for_copilot(message));
141        }
142
143        Ok(transcript)
144    }
145
146    async fn stream_from_session(
147        &self,
148        model: ResolvedCopilotModel,
149        prompt_session: PromptSession,
150    ) -> Result<LLMStream, LLMError> {
151        struct PromptCancellationGuard {
152            cancel_handle: Option<PromptSessionCancelHandle>,
153        }
154
155        impl PromptCancellationGuard {
156            fn new(cancel_handle: PromptSessionCancelHandle) -> Self {
157                Self { cancel_handle: Some(cancel_handle) }
158            }
159
160            fn disarm(&mut self) {
161                self.cancel_handle = None;
162            }
163        }
164
165        impl Drop for PromptCancellationGuard {
166            fn drop(&mut self) {
167                if let Some(cancel_handle) = self.cancel_handle.take() {
168                    cancel_handle.cancel();
169                }
170            }
171        }
172
173        let (mut updates, mut runtime_requests, completion, cancel_handle) = prompt_session.into_parts();
174        let stream = stream! {
175            let mut cancellation_guard = PromptCancellationGuard::new(cancel_handle);
176            let completion = completion;
177            tokio::pin!(completion);
178
179            let mut content = String::new();
180            let mut reasoning = String::new();
181
182            loop {
183                tokio::select! {
184                    update = updates.recv() => {
185                        match update {
186                            Some(PromptUpdate::Text(delta)) => {
187                                content.push_str(&delta);
188                                yield Ok(LLMStreamEvent::Token { delta });
189                            }
190                            Some(PromptUpdate::Thought(delta)) => {
191                                let delta = if !reasoning.is_empty()
192                                    && !reasoning.ends_with('\n')
193                                    && !delta.starts_with('\n')
194                                {
195                                    format!("\n{delta}")
196                                } else {
197                                    delta
198                                };
199                                reasoning.push_str(&delta);
200                                yield Ok(LLMStreamEvent::Reasoning { delta });
201                            }
202                            None => {}
203                        }
204                    }
205                    runtime_request = runtime_requests.recv() => {
206                        if let Some(runtime_request) = runtime_request {
207                            let response = match runtime_request {
208                                CopilotRuntimeRequest::Permission(request) => {
209                                    request.respond(crate::copilot::CopilotPermissionDecision::DeniedNoApprovalRule)
210                                }
211                                CopilotRuntimeRequest::ToolCall(request) => {
212                                    let tool_name = request.request.tool_name.clone();
213                                    request.respond(CopilotToolCallResponse::Failure(CopilotToolCallFailure {
214                                        text_result_for_llm: format!(
215                                            "GitHub Copilot tool execution is not available in this runtime. Tool `{tool_name}` was not executed."
216                                        ),
217                                        error: format!(
218                                            "tool '{tool_name}' cannot be executed outside the VT Code agent runloop session"
219                                        ),
220                                    }))
221                                }
222                                CopilotRuntimeRequest::TerminalCreate(_)
223                                | CopilotRuntimeRequest::TerminalOutput(_)
224                                | CopilotRuntimeRequest::TerminalRelease(_)
225                                | CopilotRuntimeRequest::TerminalKill(_)
226                                | CopilotRuntimeRequest::TerminalWaitForExit(_) => {
227                                    continue;
228                                }
229                                CopilotRuntimeRequest::ObservedToolCall(_) => {
230                                    continue;
231                                }
232                                CopilotRuntimeRequest::CompatibilityNotice(_) => {
233                                    continue;
234                                }
235                            };
236                            if let Err(err) = response {
237                                yield Err(map_copilot_error(err));
238                                break;
239                            }
240                        }
241                    }
242                    result = &mut completion => {
243                        let completion = match result.context("copilot acp prompt task join failed") {
244                            Ok(completion) => completion,
245                            Err(err) => {
246                                yield Err(map_copilot_error(err));
247                                break;
248                            }
249                        };
250                        let completion = match completion {
251                            Ok(completion) => completion,
252                            Err(err) => {
253                                yield Err(map_copilot_error(err));
254                                break;
255                            }
256                        };
257                        let finish_reason = map_stop_reason(&completion.stop_reason);
258                        while let Ok(update) = updates.try_recv() {
259                            match update {
260                                PromptUpdate::Text(delta) => {
261                                    content.push_str(&delta);
262                                    yield Ok(LLMStreamEvent::Token { delta });
263                                }
264                                PromptUpdate::Thought(delta) => {
265                                    let delta = if !reasoning.is_empty()
266                                        && !reasoning.ends_with('\n')
267                                        && !delta.starts_with('\n')
268                                    {
269                                        format!("\n{delta}")
270                                    } else {
271                                        delta
272                                    };
273                                    reasoning.push_str(&delta);
274                                    yield Ok(LLMStreamEvent::Reasoning { delta });
275                                }
276                            }
277                        }
278
279                        let mut response =
280                            LLMResponse::new(model.request_model.clone(), content.clone());
281                        response.finish_reason = finish_reason;
282                        if !reasoning.is_empty() {
283                            response.reasoning = Some(reasoning.clone());
284                        }
285                        cancellation_guard.disarm();
286                        yield Ok(LLMStreamEvent::Completed {
287                            response: Box::new(response),
288                        });
289                        break;
290                    }
291                }
292            }
293        };
294
295        Ok(Box::pin(stream))
296    }
297
298    async fn start_prompt_session_impl(
299        &self,
300        request: LLMRequest,
301        tools: &[ToolDefinition],
302    ) -> Result<PromptSession, LLMError> {
303        self.validate_request(&request)?;
304        let model = self.resolve_model(&request)?;
305        let transcript = self.build_transcript(&request)?;
306        let client = self.client(&model, tools).await?;
307        client.start_prompt(transcript).await.map_err(map_copilot_error)
308    }
309}
310
311#[async_trait]
312impl LLMProvider for CopilotProvider {
313    fn name(&self) -> &str {
314        COPILOT_PROVIDER_KEY
315    }
316
317    fn supports_streaming(&self) -> bool {
318        true
319    }
320
321    fn supports_non_streaming(&self, _model: &str) -> bool {
322        false
323    }
324
325    fn supports_reasoning(&self, _model: &str) -> bool {
326        true
327    }
328
329    fn supports_tools(&self, _model: &str) -> bool {
330        true
331    }
332
333    fn supports_structured_output(&self, _model: &str) -> bool {
334        false
335    }
336
337    fn supports_vision(&self, _model: &str) -> bool {
338        false
339    }
340
341    async fn generate(&self, request: LLMRequest) -> Result<LLMResponse, LLMError> {
342        let model = self.resolve_model(&request)?;
343        let mut stream = self.stream(request).await?;
344        let mut content = String::new();
345        let mut reasoning = String::new();
346        let mut completed = None;
347
348        use futures::StreamExt;
349        while let Some(event) = stream.next().await {
350            match event? {
351                LLMStreamEvent::Token { delta } => content.push_str(&delta),
352                LLMStreamEvent::Reasoning { delta } => reasoning.push_str(&delta),
353                LLMStreamEvent::ReasoningSignature { .. } => {}
354                LLMStreamEvent::ReasoningStage { .. } => {}
355                LLMStreamEvent::Completed { response } => {
356                    completed = Some(*response);
357                    break;
358                }
359            }
360        }
361
362        Ok(completed.unwrap_or_else(|| {
363            let mut response = LLMResponse::new(model.request_model.clone(), content);
364            if !reasoning.is_empty() {
365                response.reasoning = Some(reasoning);
366            }
367            response
368        }))
369    }
370
371    async fn stream(&self, request: LLMRequest) -> Result<LLMStream, LLMError> {
372        self.validate_request(&request)?;
373        let model = self.resolve_model(&request)?;
374        let transcript = self.build_transcript(&request)?;
375        let client = self.client(&model, &[]).await?;
376        let prompt_session = client.start_prompt(transcript).await.map_err(map_copilot_error)?;
377        self.stream_from_session(model, prompt_session).await
378    }
379
380    fn start_copilot_prompt_session<'a>(
381        &'a self,
382        request: LLMRequest,
383        tools: &'a [ToolDefinition],
384    ) -> Option<CopilotPromptSessionFuture<'a>> {
385        Some(Box::pin(async move { self.start_prompt_session_impl(request, tools).await }))
386    }
387
388    fn supported_models(&self) -> Vec<String> {
389        supported_models_for_provider(COPILOT_PROVIDER_KEY)
390            .map(|models| models.iter().map(|model| (*model).to_string()).collect())
391            .unwrap_or_else(|| {
392                copilot_models::SUPPORTED_MODELS
393                    .iter()
394                    .map(|model| (*model).to_string())
395                    .collect()
396            })
397    }
398
399    fn validate_request(&self, request: &LLMRequest) -> Result<(), LLMError> {
400        validate_request_common(request, "GitHub Copilot", COPILOT_PROVIDER_KEY, None)?;
401
402        if request.tools.as_ref().is_some_and(|tools| !tools.is_empty()) {
403            return Err(invalid_request("GitHub Copilot in VT Code v1 does not accept VT Code tool definitions."));
404        }
405
406        if request.output_format.is_some() {
407            return Err(invalid_request("GitHub Copilot in VT Code v1 does not support structured output."));
408        }
409
410        Ok(())
411    }
412}
413
414fn append_block(buffer: &mut String, label: &str, text: &str) {
415    if text.trim().is_empty() {
416        return;
417    }
418    if !buffer.is_empty() {
419        buffer.push_str("\n\n");
420    }
421    buffer.push_str(label);
422    buffer.push_str(":\n");
423    buffer.push_str(text.trim());
424}
425
426fn render_message_for_copilot(message: &Message) -> String {
427    let mut sections = Vec::new();
428    let text = message.content.as_text();
429    let trimmed = text.trim();
430    if !trimmed.is_empty() {
431        sections.push(trimmed.to_string());
432    }
433
434    if let Some(tool_calls) = message.tool_calls.as_ref().filter(|calls| !calls.is_empty()) {
435        let mut tool_history = String::from("[VT Code tool call history]");
436        for call in tool_calls {
437            let (tool_name, args) = call
438                .function
439                .as_ref()
440                .map(|function| (function.name.as_str(), function.arguments.trim()))
441                .unwrap_or((call.call_type.as_str(), ""));
442            if args.is_empty() {
443                let _ = write!(tool_history, "\n- {tool_name} id={}", call.id);
444            } else {
445                let _ = write!(tool_history, "\n- {tool_name} id={} args={args}", call.id);
446            }
447        }
448        sections.push(tool_history);
449    }
450
451    if message.role == MessageRole::Tool {
452        let mut tool_result = String::from("[VT Code tool result]");
453        if let Some(tool_call_id) = message.tool_call_id.as_deref() {
454            let _ = write!(tool_result, "\n- tool_call_id: {tool_call_id}");
455        }
456        if let Some(origin_tool) = message.origin_tool.as_deref() {
457            let _ = write!(tool_result, "\n- tool: {origin_tool}");
458        }
459        sections.insert(0, tool_result);
460    }
461
462    let (image_count, file_count) = count_non_text_parts(message);
463    if image_count > 0 {
464        sections.push(format!(
465            "[VT Code omitted {image_count} image input{} because GitHub Copilot v1 only accepts text input.]",
466            plural_suffix(image_count)
467        ));
468    }
469    if file_count > 0 {
470        sections.push(format!(
471            "[VT Code omitted {file_count} file attachment{} because GitHub Copilot v1 only accepts text input.]",
472            plural_suffix(file_count)
473        ));
474    }
475
476    sections.join("\n\n")
477}
478
479fn count_non_text_parts(message: &Message) -> (usize, usize) {
480    match &message.content {
481        crate::provider::MessageContent::Text(_) => (0, 0),
482        crate::provider::MessageContent::Parts(parts) => {
483            let image_count = parts.iter().filter(|part| part.is_image()).count();
484            let file_count = parts.iter().filter(|part| part.is_file()).count();
485            (image_count, file_count)
486        }
487    }
488}
489
490fn plural_suffix(count: usize) -> &'static str {
491    if count == 1 { "" } else { "s" }
492}
493
494fn invalid_request(message: &str) -> LLMError {
495    LLMError::InvalidRequest { message: message.to_string(), metadata: None }
496}
497
498fn map_copilot_error(error: anyhow::Error) -> LLMError {
499    let message = error.to_string();
500    if message.contains("rpc error -32001") || message.contains("Authentication required") {
501        return LLMError::Authentication {
502            message: "GitHub Copilot authentication is required. Run `vtcode login copilot`.".to_string(),
503            metadata: None,
504        };
505    }
506
507    LLMError::Provider { message, metadata: None }
508}
509
510fn map_stop_reason(stop_reason: &str) -> crate::provider::FinishReason {
511    match stop_reason {
512        "end_turn" => crate::provider::FinishReason::Stop,
513        "max_tokens" => crate::provider::FinishReason::Length,
514        "refusal" => crate::provider::FinishReason::Refusal,
515        "cancelled" => crate::provider::FinishReason::Error("cancelled".to_string()),
516        other => crate::provider::FinishReason::Error(other.to_string()),
517    }
518}
519
520fn copilot_tool_signature(tools: &[ToolDefinition]) -> String {
521    let mut signature_parts = tools
522        .iter()
523        .filter_map(|tool| {
524            let function = tool.function.as_ref()?;
525            Some(format!("{}:{}", function.name, serde_json::to_string(&function.parameters).ok()?))
526        })
527        .collect::<Vec<_>>();
528    signature_parts.sort_unstable();
529    signature_parts.join("|")
530}
531
532fn normalize_copilot_model_id(model: &str) -> Option<Option<String>> {
533    let trimmed = model.trim();
534    if trimmed.is_empty() {
535        return None;
536    }
537
538    match trimmed {
539        copilot_models::AUTO => Some(None),
540        copilot_models::GPT_5_CODEX => Some(Some("gpt-5.2-codex".to_string())),
541        copilot_models::GPT_5_1_CODEX_MAX => Some(Some("gpt-5.1-codex-max".to_string())),
542        copilot_models::GPT_5_6_SOL => Some(Some("gpt-5.6-sol".to_string())),
543        copilot_models::GPT_5_6_LUNA => Some(Some("gpt-5.6-luna".to_string())),
544        copilot_models::CLAUDE_SONNET_5 => Some(Some("claude-sonnet-4.6".to_string())),
545        _ if trimmed.contains(char::is_whitespace) => None,
546        _ => Some(Some(trimmed.to_string())),
547    }
548}
549
550#[cfg(test)]
551mod tests {
552    use super::CopilotProvider;
553    use super::normalize_copilot_model_id;
554    use crate::provider::{ContentPart, LLMProvider, LLMRequest, Message, ToolCall};
555    use std::path::PathBuf;
556    use std::sync::Arc;
557    use vtcode_config::constants::models::copilot as copilot_models;
558
559    fn provider() -> CopilotProvider {
560        CopilotProvider::from_config(None, None, Some(PathBuf::from("/tmp")))
561    }
562
563    #[test]
564    fn transcript_flattens_system_user_and_assistant_messages() {
565        let provider = provider();
566        let request = LLMRequest {
567            system_prompt: Some(Arc::from("Follow repository conventions.")),
568            messages: Arc::new(vec![
569                Message::user("Inspect the diff.".to_string()),
570                Message::assistant("The diff looks safe.".to_string()),
571            ]),
572            ..Default::default()
573        };
574
575        let transcript = provider.build_transcript(&request).expect("transcript should build");
576
577        assert_eq!(
578            transcript,
579            "System:\nFollow repository conventions.\n\nUser:\nInspect the diff.\n\nAssistant:\nThe diff looks safe."
580        );
581    }
582
583    #[test]
584    fn curated_model_mapping_uses_auto_as_empty_override() {
585        assert_eq!(normalize_copilot_model_id(copilot_models::AUTO), Some(None));
586        assert_eq!(normalize_copilot_model_id(copilot_models::GPT_5_6_SOL), Some(Some("gpt-5.6-sol".to_string())));
587    }
588
589    #[test]
590    fn normalize_copilot_model_id_accepts_raw_model_ids() {
591        assert_eq!(normalize_copilot_model_id("gpt-5-codex"), Some(Some("gpt-5-codex".to_string())));
592        assert_eq!(normalize_copilot_model_id("gpt 5.3"), None);
593    }
594
595    #[test]
596    fn validate_request_allows_tool_history_followups() {
597        let provider = provider();
598        let request = LLMRequest {
599            messages: Arc::new(vec![Message::tool_response("call-1".to_string(), "tool output".to_string())]),
600            ..Default::default()
601        };
602
603        provider
604            .validate_request(&request)
605            .expect("tool history should be flattened for Copilot");
606    }
607
608    #[test]
609    fn transcript_flattens_tool_history_and_image_inputs() {
610        let provider = provider();
611        let request = LLMRequest {
612            messages: Arc::new(vec![
613                Message::assistant_with_tools(
614                    "Running checks.".to_string(),
615                    vec![ToolCall::function(
616                        "call-1".to_string(),
617                        "exec_command".to_string(),
618                        r#"{"cmd":"cargo check"}"#.to_string(),
619                    )],
620                ),
621                Message::tool_response_with_origin(
622                    "call-1".to_string(),
623                    "cargo check completed successfully.".to_string(),
624                    "exec_command".to_string(),
625                ),
626                Message::user_with_parts(vec![
627                    ContentPart::text("Tell me more.".to_string()),
628                    ContentPart::image("AAAA".to_string(), "image/png".to_string()),
629                ]),
630            ]),
631            ..Default::default()
632        };
633
634        let transcript = provider
635            .build_transcript(&request)
636            .expect("transcript should flatten Copilot-incompatible history");
637
638        assert!(transcript.contains("Assistant:\nRunning checks."));
639        assert!(transcript.contains("[VT Code tool call history]"));
640        assert!(transcript.contains("- exec_command id=call-1 args={\"cmd\":\"cargo check\"}"));
641        assert!(transcript.contains("Tool:\n[VT Code tool result]"));
642        assert!(transcript.contains("- tool_call_id: call-1"));
643        assert!(transcript.contains("- tool: exec_command"));
644        assert!(transcript.contains("cargo check completed successfully."));
645        assert!(transcript.contains("User:\nTell me more."));
646        assert!(transcript.contains("omitted 1 image input"));
647    }
648
649    #[test]
650    fn supported_models_include_copilot_auto() {
651        let provider = provider();
652
653        assert!(provider.supported_models().iter().any(|model| model == copilot_models::AUTO));
654    }
655
656    #[test]
657    fn supports_reasoning_for_alias_and_live_raw_models() {
658        let provider = provider();
659
660        assert!(provider.supports_reasoning(copilot_models::AUTO));
661        assert!(provider.supports_reasoning("gpt-5-codex"));
662    }
663}