use std::sync::{Arc, Mutex};
use bytes::Bytes;
use rig::http_client::{
self, HttpClientExt, LazyBody, MultipartForm, Request, Response, StreamingResponse,
};
use crate::provider::auth::RefreshedAuth;
pub(crate) type KimiRefreshFn = Arc<dyn Fn() -> anyhow::Result<RefreshedAuth> + Send + Sync>;
struct TokenState {
bearer: String,
expires_at_ms: Option<i64>,
}
struct RefreshableToken {
state: Mutex<TokenState>,
refresher: KimiRefreshFn,
}
impl RefreshableToken {
fn bearer(&self) -> String {
let mut state = self.state.lock().unwrap_or_else(|p| p.into_inner());
if let Some(expires_at) = state.expires_at_ms {
let now = chrono::Utc::now().timestamp_millis();
if crate::auth::file_store::epoch_ms_is_expired(expires_at, now) {
match (self.refresher)() {
Ok(fresh) => {
state.bearer = fresh.bearer_token;
state.expires_at_ms = fresh.expires_at_ms;
}
Err(e) => {
tracing::warn!(
target: "dirge::provider",
error = %e,
"Kimi OAuth token expired and refresh failed; sending the stale token",
);
}
}
}
}
state.bearer.clone()
}
}
fn never_refresh() -> KimiRefreshFn {
Arc::new(|| {
anyhow::bail!("static Kimi credentials are not refreshable; this refresher is unreachable")
})
}
#[derive(Clone)]
pub(crate) struct KimiHttpClient {
inner: reqwest::Client,
token: Option<Arc<RefreshableToken>>,
identity_headers: Arc<http::HeaderMap>,
}
impl Default for KimiHttpClient {
fn default() -> Self {
Self {
inner: reqwest::Client::new(),
token: None,
identity_headers: Arc::new(http::HeaderMap::new()),
}
}
}
impl std::fmt::Debug for KimiHttpClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KimiHttpClient")
.field("bearer_token", &"<redacted>")
.finish()
}
}
impl KimiHttpClient {
pub(crate) fn new(bearer_token: String) -> Self {
Self::with_identity_headers(
Some(bearer_token),
None,
never_refresh(),
default_identity_headers(),
)
}
pub(crate) fn new_refreshable(
bearer_token: String,
expires_at_ms: Option<i64>,
refresher: KimiRefreshFn,
) -> Self {
Self::with_identity_headers(
Some(bearer_token),
expires_at_ms,
refresher,
default_identity_headers(),
)
}
#[cfg(test)]
fn with_identity_headers_for_test(
bearer_token: String,
expires_at_ms: Option<i64>,
refresher: KimiRefreshFn,
identity_headers: http::HeaderMap,
) -> Self {
Self::with_identity_headers(
Some(bearer_token),
expires_at_ms,
refresher,
identity_headers,
)
}
fn with_identity_headers(
bearer_token: Option<String>,
expires_at_ms: Option<i64>,
refresher: KimiRefreshFn,
identity_headers: http::HeaderMap,
) -> Self {
Self {
inner: reqwest::Client::new(),
token: bearer_token.map(|bearer| {
Arc::new(RefreshableToken {
state: Mutex::new(TokenState {
bearer,
expires_at_ms,
}),
refresher,
})
}),
identity_headers: Arc::new(identity_headers),
}
}
fn normalized_request<T>(&self, req: Request<T>) -> http_client::Result<Request<Bytes>>
where
T: Into<Bytes>,
{
let (mut parts, body) = req.into_parts();
if let Some(token) = &self.token
&& let Ok(value) = http::HeaderValue::from_str(&format!("Bearer {}", token.bearer()))
{
parts.headers.insert(http::header::AUTHORIZATION, value);
}
for (name, value) in self.identity_headers.iter() {
parts.headers.insert(name, value.clone());
}
let mut builder = Request::builder()
.method(parts.method)
.uri(parts.uri)
.version(parts.version);
if let Some(headers) = builder.headers_mut() {
*headers = parts.headers;
}
builder
.body(body.into())
.map_err(http_client::Error::Protocol)
}
}
fn default_identity_headers() -> http::HeaderMap {
let mut map = http::HeaderMap::new();
for (name, value) in crate::auth::kimi_device::kimi_device_headers() {
if let (Ok(name), Ok(value)) = (
http::HeaderName::try_from(name.as_str()),
http::HeaderValue::from_str(&value),
) {
map.insert(name, value);
}
}
map
}
impl HttpClientExt for KimiHttpClient {
fn send<T, U>(
&self,
req: Request<T>,
) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>> + Send + 'static
where
T: Into<Bytes>,
T: Send,
U: From<Bytes>,
U: Send + 'static,
{
let inner = self.inner.clone();
let req = self.normalized_request(req);
async move {
let req = req?;
inner.send(req).await
}
}
fn send_multipart<U>(
&self,
req: Request<MultipartForm>,
) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>> + Send + 'static
where
U: From<Bytes> + Send + 'static,
{
self.inner.send_multipart(req)
}
fn send_streaming<T>(
&self,
req: Request<T>,
) -> impl Future<Output = http_client::Result<StreamingResponse>> + Send
where
T: Into<Bytes> + Send,
{
let inner = self.inner.clone();
let req = self.normalized_request(req);
async move {
let req = req?;
inner.send_streaming(req).await
}
}
}
impl super::compressing_http::StreamingWithHeaders for KimiHttpClient {
fn send_streaming_with_headers(
&self,
req: http::Request<Bytes>,
) -> impl Future<Output = super::compressing_http::StreamingSend> + Send {
use super::compressing_http::StreamingSend;
let inner = self.inner.clone();
let req = self.normalized_request(req);
async move {
match req {
Ok(req) => inner.send_streaming_with_headers(req).await,
Err(e) => StreamingSend {
result: Err(e),
headers: None,
},
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn identity_headers() -> http::HeaderMap {
let mut map = http::HeaderMap::new();
map.insert(
http::HeaderName::from_static("x-msh-platform"),
http::HeaderValue::from_static("kimi_code_cli"),
);
map.insert(
http::HeaderName::from_static("x-msh-device-id"),
http::HeaderValue::from_static("device-1"),
);
map.insert(
http::header::USER_AGENT,
http::HeaderValue::from_static("dirge/0.0.0-test"),
);
map
}
fn client(
bearer: &str,
expires_at_ms: Option<i64>,
refresher: KimiRefreshFn,
) -> KimiHttpClient {
KimiHttpClient::with_identity_headers_for_test(
bearer.to_string(),
expires_at_ms,
refresher,
identity_headers(),
)
}
fn request(client: &KimiHttpClient, preexisting: Option<&str>) -> http::HeaderMap {
let mut builder = Request::builder()
.method("POST")
.uri("https://api.kimi.com/coding/v1/chat/completions");
if let Some(bearer) = preexisting {
builder = builder.header(http::header::AUTHORIZATION, bearer);
}
let req = builder.body(Bytes::from("{}")).unwrap();
client.normalized_request(req).unwrap().headers().clone()
}
fn authorization(headers: &http::HeaderMap) -> Option<String> {
headers
.get(http::header::AUTHORIZATION)
.map(|v| v.to_str().unwrap().to_string())
}
#[test]
fn refreshable_client_overwrites_authorization_with_refreshed_bearer() {
let refresher: KimiRefreshFn = Arc::new(|| {
Ok(RefreshedAuth {
bearer_token: "FRESH".to_string(),
expires_at_ms: Some(i64::MAX),
})
});
let client = client("STALE", Some(0), refresher);
assert_eq!(
authorization(&request(&client, Some("Bearer STALE"))).as_deref(),
Some("Bearer FRESH")
);
}
#[test]
fn refreshable_client_keeps_fresh_bearer_without_refreshing() {
let refresher: KimiRefreshFn = Arc::new(|| panic!("must not refresh a fresh token"));
let client = client("CURRENT", Some(i64::MAX), refresher);
assert_eq!(
authorization(&request(&client, Some("Bearer CURRENT"))).as_deref(),
Some("Bearer CURRENT")
);
}
#[test]
fn refresh_failure_falls_back_to_the_stale_bearer() {
let refresher: KimiRefreshFn = Arc::new(|| anyhow::bail!("network down"));
let client = client("STALE", Some(0), refresher);
assert_eq!(
authorization(&request(&client, Some("Bearer STALE"))).as_deref(),
Some("Bearer STALE")
);
}
#[test]
fn static_client_bearer_is_never_refreshed() {
let client = client("API-KEY", None, never_refresh());
assert_eq!(
authorization(&request(&client, Some("Bearer API-KEY"))).as_deref(),
Some("Bearer API-KEY")
);
}
#[test]
fn identity_headers_are_injected_on_every_request() {
let client = client("TOKEN", Some(i64::MAX), never_refresh());
let headers = request(&client, None);
assert_eq!(headers.get("x-msh-platform").unwrap(), "kimi_code_cli");
assert_eq!(headers.get("x-msh-device-id").unwrap(), "device-1");
assert_eq!(
headers.get(http::header::USER_AGENT).unwrap(),
"dirge/0.0.0-test"
);
}
#[test]
fn default_client_leaves_authorization_and_headers_untouched() {
let client = KimiHttpClient::default();
let headers = request(&client, Some("Bearer PREEXISTING"));
assert_eq!(
authorization(&headers).as_deref(),
Some("Bearer PREEXISTING")
);
assert!(headers.get("x-msh-platform").is_none());
}
#[test]
fn debug_redacts_bearer_token() {
let client = client("SUPER-SECRET", Some(i64::MAX), never_refresh());
let rendered = format!("{client:?}");
assert!(!rendered.contains("SUPER-SECRET"), "{rendered}");
assert!(rendered.contains("<redacted>"), "{rendered}");
}
}