Skip to main content

claude_codex/providers/kimi/
mod.rs

1pub mod auth;
2pub mod client;
3pub mod count_tokens;
4pub mod translate;
5
6use async_trait::async_trait;
7use axum::Json;
8use axum::response::{IntoResponse, Response};
9use http::StatusCode;
10use std::time::{SystemTime, UNIX_EPOCH};
11
12use crate::anthropic::error::json_error;
13use crate::anthropic::schema::{CountTokensResponse, MessagesRequest};
14use crate::monitor::usage_from_anthropic_sse;
15use crate::provider::{
16    CliHandlers, Generation, GenerationBody, Provider, ProviderError, ProviderErrorKind,
17    RequestContext,
18};
19use crate::providers::kimi::auth::token_store::file_store;
20use crate::providers::kimi::translate::accumulate::accumulate_response;
21use crate::providers::kimi::translate::model_allowlist::{assert_allowed_model, resolve_model};
22use crate::providers::kimi::translate::request::{TranslateOptions, translate_request};
23use crate::providers::kimi::translate::stream::translate_stream_bytes;
24use crate::registry::KIMI_MODELS;
25
26fn now_ms() -> u64 {
27    SystemTime::now()
28        .duration_since(UNIX_EPOCH)
29        .unwrap_or_default()
30        .as_millis() as u64
31}
32
33pub struct KimiProvider;
34
35impl Default for KimiProvider {
36    fn default() -> Self {
37        Self::new()
38    }
39}
40
41impl KimiProvider {
42    pub fn new() -> Self {
43        Self
44    }
45}
46
47#[async_trait]
48impl Provider for KimiProvider {
49    fn name(&self) -> &'static str {
50        "kimi"
51    }
52
53    fn supported_models(&self) -> Vec<String> {
54        KIMI_MODELS.iter().map(|s| s.to_string()).collect()
55    }
56
57    fn cli(&self) -> &'static dyn CliHandlers {
58        &KIMI_CLI
59    }
60
61    async fn handle_messages(&self, body: MessagesRequest, ctx: RequestContext) -> Response {
62        let message_id = format!("msg_{}", uuid::Uuid::new_v4().to_string().replace('-', ""));
63        let want_stream = body.stream;
64        let model = body.model.as_deref().unwrap_or("kimi-for-coding");
65        let resolved = resolve_model(model);
66
67        if let Err(e) = assert_allowed_model(&resolved) {
68            return json_error(
69                StatusCode::BAD_REQUEST,
70                "invalid_request_error",
71                format!(
72                    "Model \"{model}\" resolves to unsupported model \"{}\"",
73                    e.model
74                ),
75            );
76        }
77        if let Some(monitor) = ctx.monitor.as_ref() {
78            monitor.model_resolved(&ctx.req_id, &resolved);
79        }
80
81        let translated = match translate_request(
82            &body,
83            TranslateOptions {
84                session_id: ctx.session_id.clone(),
85            },
86        ) {
87            Ok(t) => t,
88            Err(e) => {
89                return json_error(
90                    StatusCode::BAD_REQUEST,
91                    "invalid_request_error",
92                    e.to_string(),
93                );
94            }
95        };
96
97        // KimiHttpClient uses a blocking client whose lifecycle belongs on a
98        // blocking thread.
99        if let Some(monitor) = ctx.monitor.as_ref() {
100            monitor.upstream_started(&ctx.req_id);
101        }
102        let upstream = match tokio::task::spawn_blocking(move || {
103            let client = client::KimiHttpClient::new();
104            let result = client.post_kimi(&translated);
105            drop(client);
106            result
107        })
108        .await
109        {
110            Ok(Ok(r)) => r,
111            Ok(Err(e)) => {
112                return map_kimi_error_to_response(&e);
113            }
114            Err(join_err) => {
115                return json_error(
116                    StatusCode::BAD_GATEWAY,
117                    "api_error",
118                    format!("Blocking task join error: {join_err}"),
119                );
120            }
121        };
122
123        if want_stream {
124            let sse_bytes = match translate_stream_bytes(&upstream.body, &message_id, model) {
125                Ok(b) => b,
126                Err(e) => {
127                    return json_error(
128                        StatusCode::BAD_GATEWAY,
129                        "api_error",
130                        format!("Stream translation error: {e}"),
131                    );
132                }
133            };
134            if let Some(monitor) = ctx.monitor.as_ref() {
135                let (input_tokens, output_tokens) = usage_from_anthropic_sse(&sse_bytes);
136                monitor.stream_progress(
137                    &ctx.req_id,
138                    sse_bytes.len() as u64,
139                    count_sse_events(&sse_bytes),
140                    input_tokens,
141                    output_tokens,
142                );
143            }
144
145            let headers = [
146                (http::header::CONTENT_TYPE, "text/event-stream"),
147                (http::header::CACHE_CONTROL, "no-cache"),
148                (http::header::CONNECTION, "keep-alive"),
149            ];
150            (headers, sse_bytes).into_response()
151        } else {
152            match accumulate_response(&upstream.body, &message_id, model) {
153                Ok(json) => {
154                    if let Some(monitor) = ctx.monitor.as_ref() {
155                        monitor.usage_updated(
156                            &ctx.req_id,
157                            json.pointer("/usage/input_tokens").and_then(|v| v.as_u64()),
158                            json.pointer("/usage/output_tokens")
159                                .and_then(|v| v.as_u64()),
160                        );
161                    }
162                    (StatusCode::OK, Json(json)).into_response()
163                }
164                Err(e) => json_error(
165                    StatusCode::BAD_GATEWAY,
166                    "api_error",
167                    format!("Accumulation error: {e}"),
168                ),
169            }
170        }
171    }
172
173    async fn handle_count_tokens(&self, body: MessagesRequest, ctx: RequestContext) -> Response {
174        let model = body.model.as_deref().unwrap_or("kimi-for-coding");
175        let resolved = resolve_model(model);
176        if let Some(monitor) = ctx.monitor.as_ref() {
177            monitor.model_resolved(&ctx.req_id, &resolved);
178        }
179        let tokens = count_tokens::count_tokens(&body);
180        if let Some(monitor) = ctx.monitor.as_ref() {
181            monitor.usage_updated(&ctx.req_id, Some(tokens), None);
182        }
183        (
184            StatusCode::OK,
185            Json(CountTokensResponse {
186                input_tokens: tokens,
187            }),
188        )
189            .into_response()
190    }
191
192    async fn generate_anthropic_stream(
193        &self,
194        mut body: MessagesRequest,
195        ctx: RequestContext,
196    ) -> Result<Generation, ProviderError> {
197        body.stream = true;
198        let requested = body
199            .model
200            .clone()
201            .unwrap_or_else(|| "kimi-for-coding".to_string());
202        let resolved = resolve_model(&requested);
203        assert_allowed_model(&resolved).map_err(|error| {
204            ProviderError::new(
205                StatusCode::BAD_REQUEST,
206                ProviderErrorKind::InvalidRequest,
207                format!(
208                    "Model \"{requested}\" resolves to unsupported model \"{}\"",
209                    error.model
210                ),
211            )
212        })?;
213        if let Some(monitor) = ctx.monitor.as_ref() {
214            monitor.model_resolved(&ctx.req_id, &resolved);
215        }
216        let translated = translate_request(
217            &body,
218            TranslateOptions {
219                session_id: ctx.session_id.clone(),
220            },
221        )
222        .map_err(|error| {
223            ProviderError::new(
224                StatusCode::BAD_REQUEST,
225                ProviderErrorKind::InvalidRequest,
226                error.to_string(),
227            )
228        })?;
229        if let Some(traffic) = ctx.traffic.as_ref() {
230            traffic.write_json(
231                "020-upstream-request",
232                &serde_json::to_value(&translated).unwrap_or_default(),
233            );
234        }
235        if let Some(monitor) = ctx.monitor.as_ref() {
236            monitor.upstream_started(&ctx.req_id);
237        }
238        let upstream = tokio::task::spawn_blocking(move || {
239            let client = client::KimiHttpClient::new();
240            let result = client.post_kimi(&translated);
241            drop(client);
242            result
243        })
244        .await
245        .map_err(|error| {
246            ProviderError::new(
247                StatusCode::BAD_GATEWAY,
248                ProviderErrorKind::Api,
249                format!("Blocking task join error: {error}"),
250            )
251        })?
252        .map_err(kimi_provider_error)?;
253        if let Some(traffic) = ctx.traffic.as_ref() {
254            traffic.write_bytes("032-upstream-response-body.sse", &upstream.body);
255        }
256        let message_id = format!("msg_{}", uuid::Uuid::new_v4().simple());
257        let sse =
258            translate_stream_bytes(&upstream.body, &message_id, &requested).map_err(|error| {
259                ProviderError::new(
260                    StatusCode::BAD_GATEWAY,
261                    ProviderErrorKind::Api,
262                    format!("Stream translation error: {error}"),
263                )
264            })?;
265        if let Some(traffic) = ctx.traffic.as_ref() {
266            traffic.write_bytes("050-anthropic-intermediate.sse", &sse);
267        }
268        if let Some(monitor) = ctx.monitor.as_ref() {
269            let (input_tokens, output_tokens) = usage_from_anthropic_sse(&sse);
270            monitor.stream_progress(
271                &ctx.req_id,
272                sse.len() as u64,
273                count_sse_events(&sse),
274                input_tokens,
275                output_tokens,
276            );
277        }
278        Ok(Generation {
279            body: GenerationBody::BufferedSse(sse.into()),
280            resolved_model: resolved,
281        })
282    }
283}
284
285fn count_sse_events(bytes: &[u8]) -> u64 {
286    String::from_utf8_lossy(bytes).matches("event:").count() as u64
287}
288
289fn kimi_provider_error(err: client::KimiError) -> ProviderError {
290    let (status, kind) = match err.status {
291        401 | 403 => (StatusCode::UNAUTHORIZED, ProviderErrorKind::Authentication),
292        429 => (StatusCode::TOO_MANY_REQUESTS, ProviderErrorKind::RateLimit),
293        _ => (StatusCode::BAD_GATEWAY, ProviderErrorKind::Api),
294    };
295    let mut error = ProviderError::new(status, kind, err.detail.unwrap_or(err.message));
296    if err.status == 429 {
297        error.retry_after = Some(err.retry_after.unwrap_or_else(|| "5".to_string()));
298    }
299    error
300}
301
302fn map_kimi_error_to_response(err: &client::KimiError) -> Response {
303    match err.status {
304        401 | 403 => json_error(
305            StatusCode::UNAUTHORIZED,
306            "authentication_error",
307            err.detail.as_deref().unwrap_or("Authentication failed"),
308        ),
309        429 => {
310            let retry_after = err.retry_after.as_deref().unwrap_or("5");
311            let resp = json_error(
312                StatusCode::TOO_MANY_REQUESTS,
313                "rate_limit_error",
314                &err.message,
315            );
316            // Forward retry-after header
317            let headers = [(http::header::RETRY_AFTER, retry_after)];
318            (headers, resp).into_response()
319        }
320        _ => json_error(
321            StatusCode::BAD_GATEWAY,
322            "api_error",
323            err.detail.as_deref().unwrap_or("Upstream error"),
324        ),
325    }
326}
327
328// ---------------------------------------------------------------------------
329// CLI
330// ---------------------------------------------------------------------------
331
332pub(crate) struct KimiCli;
333
334impl CliHandlers for KimiCli {
335    fn login(&self) -> Result<(), anyhow::Error> {
336        let tokens = auth::login::run_device_login()?;
337        let store = file_store();
338        let manager = auth::manager::KimiAuthManager::new(store);
339        let saved = manager.persist_initial_tokens(&tokens)?;
340        println!("Auth saved in {}", manager.store.auth_path());
341        if let Some(ref uid) = saved.user_id {
342            println!("User: {uid}");
343        }
344        println!("Authentication complete");
345        Ok(())
346    }
347
348    fn device(&self) -> Result<(), anyhow::Error> {
349        self.login()
350    }
351
352    fn status(&self) -> Result<(), anyhow::Error> {
353        let store = file_store();
354        let stored = store.load_auth()?;
355        match stored {
356            Some(auth) => {
357                println!("Auth path: {}", store.auth_path());
358                println!("Authenticated: true");
359                if let Some(ref uid) = auth.user_id {
360                    println!("User: {uid}");
361                }
362                if let Some(ref scope) = auth.scope {
363                    println!("Scope: {scope}");
364                }
365                let remaining = auth.expires.saturating_sub(now_ms()) / 1000;
366                println!("Expires in {remaining}s");
367                Ok(())
368            }
369            None => {
370                anyhow::bail!("Not authenticated");
371            }
372        }
373    }
374
375    fn logout(&self) -> Result<(), anyhow::Error> {
376        let store = file_store();
377        store.clear_auth()?;
378        println!("Logged out");
379        Ok(())
380    }
381}
382
383pub(crate) static KIMI_CLI: KimiCli = KimiCli;