rig_core/providers/copilot/
wire.rs1use crate::error::ProviderError;
16use crate::wire::Flow;
17use serde::{Deserialize, Serialize};
18
19use crate::client::env::{self, EnvError};
20use crate::completion::CompletionRequest;
21use crate::error::EncodeError;
22use crate::model::{ModelInfo, ModelList};
23use crate::operation::{Completion, ModelListing, ModelPage};
24use crate::providers::internal::wire::classify_untyped_line;
25use crate::providers::openai::responses_api::SystemInstructionsPlacement;
26pub use crate::providers::openai::wire::Embeddings;
30use crate::providers::openai::wire::{
31 Dialect, DialectHooks, EmbeddingQuirks, OpenAIConfig, OpenAiDecoder, OpenAiWire, Quirks,
32 ResponsesQuirks, Route,
33};
34use crate::wire::{
35 Body, Decoder, Descriptor, Encoded, Framing, Mode, Out, Secret, Wire, WireEvent, WireFrame,
36};
37
38use super::{CopilotIntent, PROVIDER_NAME};
39
40const REQUEST_ID_HEADER: Option<&str> = Some("x-request-id");
43
44const PRIMARY_API_KEY_ENV: &str = "GITHUB_COPILOT_API_KEY";
46
47const API_KEY_ENV: [&str; 2] = ["GITHUB_COPILOT_API_KEY", "COPILOT_API_KEY"];
49
50const BASE_URL_ENV: &[&str] = &["GITHUB_COPILOT_API_BASE", "COPILOT_BASE_URL"];
52
53pub const DIALECT: Dialect = Dialect {
57 base_url_env: Some("GITHUB_COPILOT_API_BASE"),
58 request_id_header: REQUEST_ID_HEADER,
59 quirks: Quirks {
60 hooks: Some(&HOOKS),
61 verify_path: "",
62 base_url_env_alias: Some("COPILOT_BASE_URL"),
63 embedding: EmbeddingQuirks {
64 requires_usage: false,
65 ..EmbeddingQuirks::openai()
66 },
67 responses: ResponsesQuirks {
68 strict_tools_by_default: true,
69 system_instructions: SystemInstructionsPlacement::InputSystemMessages,
70 ..ResponsesQuirks::openai()
71 },
72 ..Quirks::openai()
73 },
74 ..Dialect::gateway(
75 PROVIDER_NAME,
76 super::GITHUB_COPILOT_API_BASE_URL,
77 "GITHUB_COPILOT_API_KEY",
78 )
79};
80
81static HOOKS: DialectHooks = DialectHooks {
82 default_endpoint: Some(super::base_url_from_token),
83 model_route: Some(|model| {
84 if routes_through_responses(model) {
85 Route::Responses
86 } else {
87 Route::Chat
88 }
89 }),
90 completion_envelope: Some(|provider, request, builder| {
91 completion_envelope(provider, request, builder, CopilotIntent::default())
92 }),
93 modality_envelope: Some(|provider, request| {
95 stamp(
96 request,
97 provider.api_key.expose(),
98 "user",
99 false,
100 CopilotIntent::Panel,
101 )
102 }),
103};
104
105fn completion_envelope(
107 provider: &OpenAIConfig,
108 request: &CompletionRequest,
109 mut builder: http::request::Builder,
110 intent: CopilotIntent,
111) -> http::request::Builder {
112 for (name, value) in super::default_headers(
113 provider.api_key.expose(),
114 super::request_initiator(request),
115 super::request_has_vision(request),
116 intent,
117 ) {
118 if let Some(headers) = builder.headers_mut() {
119 headers.remove(name);
120 }
121 builder = builder.header(name, value);
122 }
123 builder
124}
125
126pub fn routes_through_responses(model: &str) -> bool {
128 model.to_ascii_lowercase().contains("codex")
129}
130
131#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
136pub struct CopilotConfig {
137 pub api_key: Secret,
139 pub base_url: String,
141}
142
143impl CopilotConfig {
144 pub fn new(api_key: impl Into<Secret>) -> Self {
148 credential_of(&OpenAIConfig::with_key(&DIALECT, api_key))
149 }
150
151 pub fn from_auth(context: &super::auth::AuthContext) -> Self {
154 let mut provider = Self::new(context.api_key.clone());
155 if let Some(api_base) = &context.api_base {
156 provider.base_url = api_base.clone();
157 }
158 provider
159 }
160
161 pub fn from_env() -> Result<Self, EnvError> {
166 let Some(api_key) = first_env(&API_KEY_ENV)? else {
167 return Err(EnvError::Variable {
168 name: PRIMARY_API_KEY_ENV,
169 source: std::env::VarError::NotPresent,
170 });
171 };
172 let mut provider = Self::new(api_key);
173 if let Some(base_url) = first_env(BASE_URL_ENV)? {
174 provider.base_url = base_url;
175 }
176 Ok(provider)
177 }
178
179 pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
181 self.base_url = base_url.into();
182 self
183 }
184
185 pub(crate) fn completion(&self, model: impl Into<String>) -> CopilotWire {
187 self.wire_for(model)
188 }
189
190 fn wire_for(&self, model: impl Into<String>) -> CopilotWire {
192 CopilotWire {
193 wire: self.openai().completion(model),
194 intent: CopilotIntent::default(),
195 }
196 }
197
198 pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
201 Embeddings::new(self.openai(), model, ndims)
202 }
203
204 pub(crate) fn models(&self) -> Models {
206 Models {
207 provider: self.clone(),
208 }
209 }
210
211 fn openai(&self) -> OpenAIConfig {
213 OpenAIConfig::with_key(&DIALECT, self.api_key.clone()).with_base_url(self.base_url.clone())
214 }
215
216 fn uri(&self, path: &str) -> String {
218 format!("{}{path}", self.base_url.trim_end_matches('/'))
219 }
220}
221
222fn first_env(names: &[&'static str]) -> Result<Option<String>, EnvError> {
225 for name in names {
226 if let Some(value) = env::optional(name)?.filter(|value| !value.trim().is_empty()) {
227 return Ok(Some(value));
228 }
229 }
230 Ok(None)
231}
232
233fn stamp(
239 request: &mut http::Request<Body>,
240 api_key: &str,
241 initiator: &'static str,
242 has_vision: bool,
243 intent: CopilotIntent,
244) -> Result<(), http::Error> {
245 let map = request.headers_mut();
246 for (name, value) in super::default_headers(api_key, initiator, has_vision, intent) {
247 map.insert(
248 http::HeaderName::from_bytes(name.as_bytes())?,
249 http::HeaderValue::from_str(&value)?,
250 );
251 }
252 Ok(())
253}
254
255#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
259pub struct CopilotWire {
260 pub wire: OpenAiWire,
262 pub intent: CopilotIntent,
264}
265
266impl CopilotWire {
267 pub fn intent(&self) -> CopilotIntent {
269 self.intent
270 }
271
272 pub fn with_intent(mut self, intent: CopilotIntent) -> Self {
274 self.intent = intent;
275 self
276 }
277
278 pub fn with_panel_intent(self) -> Self {
280 self.with_intent(CopilotIntent::Panel)
281 }
282
283 pub fn with_edits_intent(self) -> Self {
285 self.with_intent(CopilotIntent::Edits)
286 }
287
288 pub fn with_strict_tools(mut self) -> Self {
293 self.wire = self.wire.with_strict_tools();
294 self
295 }
296
297 pub fn with_tool_result_array_content(mut self) -> Self {
302 if let OpenAiWire::Chat(wire) = self.wire {
303 self.wire = OpenAiWire::Chat(wire.with_tool_result_array_content());
304 }
305 self
306 }
307}
308
309fn credential_of(provider: &OpenAIConfig) -> CopilotConfig {
312 CopilotConfig {
313 api_key: provider.api_key.clone(),
314 base_url: provider.base_url.clone(),
315 }
316}
317
318impl Wire for CopilotWire {
319 type Op = Completion;
320 type Payload = crate::wire::Encoded;
321 type Frame = crate::wire::WireFrame;
322 type Decoder<'id> = OpenAiDecoder<'id>;
323
324 fn describe(&self) -> Descriptor<'_> {
325 self.wire.describe()
326 }
327
328 fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
329 self.wire
330 .encode_with_headers(request, mode, |provider, request, builder| {
331 completion_envelope(provider, request, provider.headers(builder), self.intent)
332 })
333 }
334
335 fn decoder<'id>(&self) -> Self::Decoder<'id> {
336 self.wire.decoder()
337 }
338}
339
340#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
345pub struct Models {
346 pub provider: CopilotConfig,
348}
349
350#[derive(Debug, Deserialize)]
352pub struct ModelEntry {
353 id: String,
354 #[serde(default)]
355 name: Option<String>,
356 #[serde(default)]
357 vendor: Option<String>,
358 #[serde(default)]
359 capabilities: Option<ModelEntryCapabilities>,
360}
361
362#[derive(Debug, Deserialize)]
363struct ModelEntryCapabilities {
364 #[serde(default, rename = "type")]
365 kind: Option<String>,
366}
367
368#[derive(Debug, Deserialize)]
370pub struct ModelsReply {
371 #[serde(default)]
372 data: Vec<ModelEntry>,
373}
374
375impl ModelsReply {
376 pub fn into_models(self) -> Vec<ModelInfo> {
378 self.data.into_iter().map(ModelInfo::from).collect()
379 }
380}
381
382impl From<ModelEntry> for ModelInfo {
383 fn from(entry: ModelEntry) -> Self {
384 let mut model = ModelInfo::from_id(entry.id);
385 model.name = entry.name;
386 model.owned_by = entry.vendor;
387 if let Some(capabilities) = entry.capabilities {
388 model.r#type = capabilities.kind;
389 }
390 model
391 }
392}
393
394#[derive(Default)]
396pub struct ModelsDecoder;
397
398impl<'id> Decoder<'id, ModelListing> for ModelsDecoder {
399 type Event = ModelsReply;
400
401 fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
402 classify_untyped_line(frame.as_str().as_bytes())
403 }
404
405 fn decode(
406 &mut self,
407 event: Self::Event,
408 out: Out<'id, ModelListing>,
409 ) -> Result<Flow, ProviderError> {
410 Ok(out.end(ModelPage {
411 models: ModelList::new(event.into_models()),
412 next: None,
413 }))
414 }
415}
416
417impl Wire for Models {
418 type Op = ModelListing;
419 type Payload = crate::wire::Encoded;
420 type Frame = crate::wire::WireFrame;
421 type Decoder<'id> = ModelsDecoder;
422
423 fn describe(&self) -> Descriptor<'_> {
424 Descriptor::new(PROVIDER_NAME)
425 }
426
427 fn encode(&self, _cursor: Option<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
428 let mut request = http::Request::get(self.provider.uri(super::MODEL_LISTING_PATH))
429 .header(http::header::CONTENT_TYPE, "application/json")
430 .body(Body::empty())?;
431 stamp(
432 &mut request,
433 self.provider.api_key.expose(),
434 "user",
435 false,
436 CopilotIntent::Panel,
437 )?;
438 Ok(Encoded::new(request, Framing::Whole).with_request_id_header(REQUEST_ID_HEADER))
439 }
440
441 fn decoder<'id>(&self) -> Self::Decoder<'id> {
442 ModelsDecoder
443 }
444}
445
446#[cfg(test)]
447mod tests;