use std::fmt;
pub const TOKEN_URL : &str = "https://platform.claude.com/v1/oauth/token";
pub const CLIENT_ID : &str = "9d1c250a-e61b-44d9-88ed-5944d1962f5e";
#[ derive( Debug ) ]
pub struct TokenRefreshResult
{
pub access_token : String,
pub refresh_token : String,
pub expires_at_ms : u64,
}
#[ derive( Debug ) ]
pub enum AuthError
{
HttpTransport( String ),
ResponseParse( String ),
RateLimited,
}
impl fmt::Display for AuthError
{
#[ inline ]
fn fmt( &self, f : &mut fmt::Formatter< '_ > ) -> fmt::Result
{
match self
{
Self::HttpTransport( msg ) => write!( f, "HTTP transport error: {msg}" ),
Self::ResponseParse( field ) =>
write!( f, "response parse error: missing or malformed field '{field}'" ),
Self::RateLimited => write!( f, "rate limited (429): back off before retrying" ),
}
}
}
impl std::error::Error for AuthError {}
#[ inline ]
pub fn parse_response( body : &str, now_ms : u64 ) -> Result< TokenRefreshResult, AuthError >
{
let access_token = parse_string_field( body, "access_token" )?;
let refresh_token = parse_string_field( body, "refresh_token" )?;
let expires_in = parse_u64_field( body, "expires_in" )?;
Ok
(
TokenRefreshResult
{
access_token,
refresh_token,
expires_at_ms : now_ms + expires_in * 1000,
}
)
}
fn parse_string_field( body : &str, key : &str ) -> Result< String, AuthError >
{
let needle = format!( "\"{key}\":" );
let after_key = body
.find( needle.as_str() )
.map( | pos | &body[ pos + needle.len() .. ] )
.ok_or_else( || AuthError::ResponseParse( key.to_string() ) )?;
let after_colon = after_key.trim_start();
if !after_colon.starts_with( '"' )
{
return Err( AuthError::ResponseParse( key.to_string() ) );
}
let inner = &after_colon[ 1 .. ];
let end = inner
.find( '"' )
.ok_or_else( || AuthError::ResponseParse( key.to_string() ) )?;
Ok( inner[ .. end ].to_string() )
}
fn parse_u64_field( body : &str, key : &str ) -> Result< u64, AuthError >
{
let needle = format!( "\"{key}\":" );
let after_key = body
.find( needle.as_str() )
.map( | pos | &body[ pos + needle.len() .. ] )
.ok_or_else( || AuthError::ResponseParse( key.to_string() ) )?;
let after_colon = after_key.trim_start();
if after_colon.starts_with( '"' )
{
return Err( AuthError::ResponseParse( key.to_string() ) );
}
let digits : &str = after_colon
.find( | c : char | !c.is_ascii_digit() )
.map_or( after_colon, | end | &after_colon[ .. end ] );
digits
.parse::< u64 >()
.map_err( | _ | AuthError::ResponseParse( key.to_string() ) )
}
#[ cfg( feature = "enabled" ) ]
#[ inline ]
pub fn refresh_token( refresh_tok : &str, scope : &str ) -> Result< TokenRefreshResult, AuthError >
{
use std::time::{ SystemTime, UNIX_EPOCH };
let body = format!(
r#"{{"grant_type":"refresh_token","refresh_token":"{refresh_tok}","client_id":"{CLIENT_ID}","scope":"{scope}"}}"#
);
let config = ureq::Agent::config_builder()
.http_status_as_error( false )
.build();
let agent = ureq::Agent::new_with_config( config );
let mut resp = agent
.post( TOKEN_URL )
.header( "Content-Type", "application/json" )
.send( body.as_str() )
.map_err( |e| AuthError::HttpTransport( e.to_string() ) )?;
let status = resp.status().as_u16();
if status == 429
{
return Err( AuthError::RateLimited );
}
if status >= 400
{
return Err( AuthError::HttpTransport( format!( "HTTP {status}" ) ) );
}
let text = resp
.body_mut()
.read_to_string()
.map_err( |e| AuthError::HttpTransport( e.to_string() ) )?;
let now_ms = SystemTime::now()
.duration_since( UNIX_EPOCH )
.map_or( 0, | d | d.as_secs() * 1000 );
parse_response( &text, now_ms )
}