use super::{Api, ApiError, Owner, Renewed};
use crate::context::Context;
use crate::provider::ProviderError;
use crate::provider::codex::api::{Fresh, OpenAi};
use crate::usage::Snapshot;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Trouble {
Unauthorized,
RateLimited,
RateLimitedFor(i64),
Offline,
Server(u16),
InvalidGrant,
}
impl From<Trouble> for ApiError {
fn from(t: Trouble) -> ApiError {
match t {
Trouble::Unauthorized => ApiError::Unauthorized,
Trouble::RateLimited => ApiError::RateLimited { retry_after: None },
Trouble::RateLimitedFor(seconds) => ApiError::RateLimited {
retry_after: Some(seconds),
},
Trouble::Offline => ApiError::Network("no route to host".into()),
Trouble::Server(status) => ApiError::Unexpected { status },
Trouble::InvalidGrant => ApiError::InvalidGrant,
}
}
}
#[derive(Debug, Clone)]
pub enum Answer<T> {
Give(T),
Fail(Trouble),
}
impl<T: Clone> Answer<T> {
fn take(&self) -> Result<T, ApiError> {
match self {
Answer::Give(value) => Ok(value.clone()),
Answer::Fail(trouble) => Err((*trouble).into()),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Asked {
Owner(String),
Usage(String),
Renew(String),
}
#[derive(Debug, Default)]
struct Script {
owners: HashMap<String, Answer<Owner>>,
usage: HashMap<String, Answer<Snapshot>>,
renewals: HashMap<String, Answer<Renewed>>,
codex_renewals: HashMap<String, Answer<Fresh>>,
asked: Vec<Asked>,
}
#[derive(Debug, Default)]
pub struct ScriptedApi(Mutex<Script>);
impl ScriptedApi {
pub fn new() -> Arc<ScriptedApi> {
Arc::new(ScriptedApi::default())
}
fn script(&self) -> std::sync::MutexGuard<'_, Script> {
self.0.lock().expect("a poisoned script is a failed test")
}
pub fn owned_by(&self, access_token: &str, owner: Owner) -> &ScriptedApi {
self.script()
.owners
.insert(access_token.into(), Answer::Give(owner));
self
}
pub fn using(&self, access_token: &str, snapshot: Snapshot) -> &ScriptedApi {
self.script()
.usage
.insert(access_token.into(), Answer::Give(snapshot));
self
}
pub fn renews(&self, refresh_token: &str, renewed: Renewed) -> &ScriptedApi {
self.script()
.renewals
.insert(refresh_token.into(), Answer::Give(renewed));
self
}
pub fn token_trouble(&self, access_token: &str, trouble: Trouble) -> &ScriptedApi {
let mut script = self.script();
script
.owners
.insert(access_token.into(), Answer::Fail(trouble));
script
.usage
.insert(access_token.into(), Answer::Fail(trouble));
drop(script);
self
}
pub fn codex_renews(&self, refresh_token: &str, fresh: Fresh) -> &ScriptedApi {
self.script()
.codex_renewals
.insert(refresh_token.into(), Answer::Give(fresh));
self
}
pub fn codex_renew_trouble(&self, refresh_token: &str, trouble: Trouble) -> &ScriptedApi {
self.script()
.codex_renewals
.insert(refresh_token.into(), Answer::Fail(trouble));
self
}
pub fn renew_trouble(&self, refresh_token: &str, trouble: Trouble) -> &ScriptedApi {
self.script()
.renewals
.insert(refresh_token.into(), Answer::Fail(trouble));
self
}
pub fn asked(&self) -> Vec<Asked> {
self.script().asked.clone()
}
pub fn calls(&self) -> usize {
self.script().asked.len()
}
fn answer<T: Clone>(
&self,
asked: Asked,
pick: impl FnOnce(&Script) -> Option<&Answer<T>>,
) -> Result<T, ApiError> {
let mut script = self.script();
script.asked.push(asked);
match pick(&script) {
Some(answer) => answer.take(),
None => Err(ApiError::Unauthorized),
}
}
}
impl Api for ScriptedApi {
fn owner(&self, _ctx: &Context, access_token: &str) -> Result<Owner, ApiError> {
self.answer(Asked::Owner(access_token.into()), |s| {
s.owners.get(access_token)
})
}
fn usage(&self, _ctx: &Context, access_token: &str) -> Result<Snapshot, ApiError> {
self.answer(Asked::Usage(access_token.into()), |s| {
s.usage.get(access_token)
})
}
fn renew(
&self,
_ctx: &Context,
refresh_token: &str,
_scopes: &[String],
_client_id: Option<&str>,
) -> Result<Renewed, ApiError> {
self.answer(Asked::Renew(refresh_token.into()), |s| {
s.renewals.get(refresh_token)
})
}
}
fn from_trouble(trouble: Trouble) -> ProviderError {
let service = crate::provider::ProviderId::Codex.service();
match trouble {
Trouble::Unauthorized => ProviderError::Unauthorized,
Trouble::RateLimited => ProviderError::RateLimited {
service,
retry_after: None,
},
Trouble::RateLimitedFor(seconds) => ProviderError::RateLimited {
service,
retry_after: Some(seconds),
},
Trouble::Offline => ProviderError::Network {
service,
detail: "no route to host".into(),
},
Trouble::Server(status) => ProviderError::Unexpected { service, status },
Trouble::InvalidGrant => ProviderError::InvalidGrant { service },
}
}
impl OpenAi for ScriptedApi {
fn usage(
&self,
_ctx: &Context,
access_token: &str,
_account_id: &str,
_now: i64,
) -> Result<Snapshot, ProviderError> {
let mut script = self.script();
script.asked.push(Asked::Usage(access_token.into()));
match script.usage.get(access_token) {
Some(Answer::Give(snapshot)) => Ok(snapshot.clone()),
Some(Answer::Fail(trouble)) => Err(from_trouble(*trouble)),
None => Err(ProviderError::Unauthorized),
}
}
fn renew(&self, _ctx: &Context, refresh_token: &str) -> Result<Fresh, ProviderError> {
let service = crate::provider::ProviderId::Codex.service();
let mut script = self.script();
script.asked.push(Asked::Renew(refresh_token.into()));
match script.codex_renewals.get(refresh_token) {
Some(Answer::Give(fresh)) => Ok(fresh.clone()),
Some(Answer::Fail(trouble)) => Err(from_trouble(*trouble)),
None => Err(ProviderError::InvalidGrant { service }),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn owner(uuid: &str) -> Owner {
Owner {
account_uuid: uuid.into(),
email: "me@example.com".into(),
organization_uuid: "org".into(),
}
}
#[test]
fn it_answers_what_it_was_told_and_refuses_what_it_was_not() {
let api = ScriptedApi::new();
api.owned_by("live", owner("acc"));
let ctx = Context::new(std::path::PathBuf::from("/nowhere"));
assert_eq!(
api.owner(&ctx, "live").expect("scripted").account_uuid,
"acc"
);
assert!(matches!(
api.owner(&ctx, "someone else"),
Err(ApiError::Unauthorized)
));
}
#[test]
fn it_produces_the_failures_a_server_cannot() {
let api = ScriptedApi::new();
api.token_trouble("live", Trouble::RateLimited);
api.renew_trouble("stale", Trouble::InvalidGrant);
let ctx = Context::new(std::path::PathBuf::from("/nowhere"));
assert!(matches!(
api.owner(&ctx, "live"),
Err(ApiError::RateLimited { .. })
));
assert!(matches!(
Api::usage(&*api, &ctx, "live"),
Err(ApiError::RateLimited { .. })
));
assert!(matches!(
Api::renew(&*api, &ctx, "stale", &[], None),
Err(ApiError::InvalidGrant)
));
}
#[test]
fn it_remembers_what_it_was_asked() {
let api = ScriptedApi::new();
api.owned_by("live", owner("acc"));
let ctx = Context::new(std::path::PathBuf::from("/nowhere"));
let _ = api.owner(&ctx, "live");
let _ = Api::usage(&*api, &ctx, "live");
let _ = api.owner(&ctx, "live");
assert_eq!(api.calls(), 3);
assert_eq!(
api.asked(),
vec![
Asked::Owner("live".into()),
Asked::Usage("live".into()),
Asked::Owner("live".into()),
]
);
}
}