use serde::Deserialize;
use crate::types::{DarajaApi, DarajaEnvironment};
#[derive(Deserialize, Debug)]
pub struct GenerateAccessTokenResponse {
pub access_token: String,
pub expires_in: String,
}
pub struct Client<'a> {
pub consumer_key: &'a str,
pub consumer_secret: &'a str,
pub environment: DarajaEnvironment,
}
impl Default for Client<'_> {
fn default() -> Self {
Self::new()
}
}
impl DarajaApi for Client<'_> {
fn path(&self) -> &'static str {
"oauth/v1/generate?grant_type=client_credentials"
}
fn environment(&self) -> DarajaEnvironment {
self.environment
}
fn set_environment(&mut self, environment: DarajaEnvironment) {
self.environment = environment;
}
}
impl<'a> Client<'a> {
pub fn new() -> Self {
Self {
consumer_key: "",
consumer_secret: "",
environment: DarajaEnvironment::default(),
}
}
pub fn with_credentials(consumer_key: &'a str, consumer_secret: &'a str) -> Self {
Self {
consumer_key,
consumer_secret,
environment: DarajaEnvironment::default(),
}
}
pub fn production(self) -> Self {
DarajaApi::production(self)
}
pub async fn generate_access_token(
self,
) -> Result<GenerateAccessTokenResponse, reqwest::Error> {
let http_client = reqwest::Client::new();
http_client
.get(self.get_url())
.basic_auth(self.consumer_key, Some(self.consumer_secret))
.send()
.await?
.error_for_status()?
.json::<GenerateAccessTokenResponse>()
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::mpesa::test_support::{
TestConfig, assert_valid_access_token, assert_valid_expires_in, read_token_cache,
write_token_cache,
};
#[test]
fn new_creates_client_with_empty_credentials() {
let client = Client::new();
assert!(client.consumer_key.is_empty());
assert!(client.consumer_secret.is_empty());
}
#[test]
fn with_credentials_sets_consumer_key_and_secret() {
let client = Client::with_credentials("test-key".into(), "test-secret".into());
assert_eq!(client.consumer_key, "test-key");
assert_eq!(client.consumer_secret, "test-secret");
}
#[tokio::test]
async fn generate_access_token_fails_with_invalid_credentials() {
let client = Client::with_credentials("invalid-key".into(), "invalid-secret".into());
let err = client
.generate_access_token()
.await
.expect_err("expected request to fail with invalid credentials");
assert_eq!(err.status(), Some(reqwest::StatusCode::BAD_REQUEST));
}
#[tokio::test]
async fn generate_access_token_returns_valid_response() {
let test_config = TestConfig::load();
let client =
Client::with_credentials(&test_config.consumer_key, &test_config.consumer_secret);
if let Some((cached_token, expires_at)) = read_token_cache() {
assert_valid_access_token(&cached_token);
let remaining = expires_at
.duration_since(std::time::SystemTime::now())
.expect("cached token should not be expired");
assert_valid_expires_in(remaining.as_secs());
return;
}
let response = client.generate_access_token().await.unwrap();
assert_valid_access_token(&response.access_token);
let expires_in: u64 = response
.expires_in
.parse()
.expect("expires_in should be a positive integer");
assert_valid_expires_in(expires_in);
write_token_cache(&response.access_token, expires_in);
}
}