use crate::auth::normalize_stored_key;
use crate::pool::GlobalPool;
use crate::queue::QueueStats;
use crate::rate_limit::RateLimitInfo;
use chrono::{DateTime, Utc};
use dashmap::DashMap;
use deadpool_postgres::Pool;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::{
collections::HashMap,
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
};
use tokio::sync::{Mutex, Semaphore};
#[derive(Clone)]
#[expect(
dead_code,
reason = "legacy global queue/pool fields retained for API compatibility"
)]
pub struct AppState {
pub instances: Arc<tokio::sync::RwLock<HashMap<String, PostgresInstance>>>,
pub connections: Arc<DashMap<String, Pool>>,
pub accounts: Arc<tokio::sync::RwLock<HashMap<String, Account>>>,
pub rate_limiter: Arc<Mutex<HashMap<String, RateLimitInfo>>>,
pub connection_limits: Arc<DashMap<String, Arc<Semaphore>>>, pub transaction_sessions: Arc<DashMap<uuid::Uuid, Arc<TransactionSession>>>,
pub global_pool: Arc<GlobalPool>,
pub job_queue: Option<tokio::sync::mpsc::Sender<crate::queue::QueuedRequest>>,
pub queue_stats: Arc<QueueStats>, }
impl AppState {
pub async fn new() -> anyhow::Result<Self> {
let pool_config = crate::config::load_pool_config().await?;
let global_pool = Arc::new(GlobalPool::new(pool_config).await?);
Ok(Self {
instances: Arc::new(tokio::sync::RwLock::new(load_instances().await?)),
connections: Arc::new(DashMap::new()),
accounts: Arc::new(tokio::sync::RwLock::new(load_accounts().await?)),
rate_limiter: Arc::new(Mutex::new(HashMap::new())),
connection_limits: Arc::new(DashMap::new()),
transaction_sessions: Arc::new(DashMap::new()),
global_pool,
job_queue: None,
queue_stats: Arc::new(QueueStats::default()),
})
}
}
pub struct TransactionSession {
pub account_id: String,
pub database: String,
pub client: Arc<Mutex<Option<deadpool_postgres::Client>>>,
pub last_used: AtomicU64,
}
impl TransactionSession {
pub fn touch(&self) {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
self.last_used.store(now, Ordering::Relaxed);
}
pub async fn finish(&self, statement: &str) -> Result<(), String> {
let client = { self.client.lock().await.take() };
let Some(client) = client else {
return Err("PG-API session already finished".to_string());
};
match client.batch_execute(statement).await {
Ok(()) => {
drop(client);
Ok(())
}
Err(error) => {
let message = crate::database::postgres_error_message(&error);
drop(deadpool_postgres::Client::take(client));
Err(message)
}
}
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct PostgresInstance {
pub id: String,
pub name: String,
pub host: String,
pub port: u16,
pub superuser: String,
pub superuser_password: String,
pub instance_type: InstanceType,
pub created_at: DateTime<Utc>,
pub status: InstanceStatus,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum InstanceType {
Single,
Primary,
Replica,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum InstanceStatus {
Active,
Maintenance,
Degraded,
Offline,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct Account {
pub id: String,
pub name: String,
pub api_key: String,
pub instance_id: String,
pub databases: Vec<DatabaseAccess>,
pub role: AccountRole,
pub created_at: DateTime<Utc>,
pub last_used: DateTime<Utc>,
pub rate_limit: u32,
pub max_connections: u32,
#[serde(skip_serializing_if = "Option::is_none", default)]
pub notes: Option<String>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct DatabaseAccess {
pub database: String,
pub username: String,
pub password: String,
pub permissions: Vec<Permission>,
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
#[serde(rename_all = "UPPERCASE")]
pub enum Permission {
Select,
Insert,
Update,
Delete,
Create,
Drop,
Truncate,
References,
Trigger,
Execute,
Usage,
#[serde(rename = "CREATE_DATABASE")]
CreateDatabase,
#[serde(rename = "DROP_OWNED_DATABASE")]
DropOwnedDatabase,
Export,
Import,
All,
}
impl Account {
pub fn has_permission(&self, permission: Permission) -> bool {
self.role == AccountRole::Superuser
|| self.databases.iter().any(|database| {
database.permissions.contains(&Permission::All)
|| database.permissions.contains(&permission)
})
}
pub fn owns_database_role(&self, username: &str) -> bool {
self.databases
.iter()
.any(|database| database.username == username)
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
#[serde(rename_all = "lowercase")]
pub enum AccountRole {
#[serde(alias = "owner")]
Superuser,
#[serde(alias = "administrator")]
Admin,
Developer,
#[serde(alias = "app", alias = "readwrite")]
Application,
#[serde(alias = "read_only", alias = "read-only")]
Readonly,
}
#[derive(Debug, Deserialize)]
pub struct QueryRequest {
pub query: String,
pub database: String,
#[serde(default)]
pub params: Vec<Value>,
#[serde(default)]
#[allow(dead_code)]
pub options: QueryOptions,
}
#[derive(Debug, Deserialize, Default)]
pub struct QueryOptions {
#[serde(default)]
#[allow(dead_code)]
pub timeout_ms: Option<u64>,
#[serde(default)]
#[allow(dead_code)]
pub read_only: bool,
#[serde(default)]
#[allow(dead_code)]
pub as_transaction: bool,
}
#[derive(Debug, Serialize)]
pub struct ApiResponse<T> {
pub success: bool,
pub data: Option<T>,
pub error: Option<ErrorInfo>,
pub metadata: ResponseMetadata,
}
#[derive(Debug, Serialize)]
pub struct ErrorInfo {
pub code: String,
pub message: String,
pub details: Option<Value>,
}
#[derive(Debug, Serialize)]
pub struct ResponseMetadata {
pub request_id: String,
pub execution_time_ms: u128,
#[serde(skip_serializing_if = "Option::is_none")]
pub rows_affected: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub instance_id: Option<String>,
pub timestamp: DateTime<Utc>,
}
#[derive(Debug, Serialize)]
pub struct QueryResult {
pub rows: Vec<Value>,
pub fields: Vec<FieldMetadata>,
#[serde(skip_serializing_if = "Option::is_none")]
pub query_plan: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rows_affected: Option<u64>,
}
#[derive(Debug, Serialize)]
pub struct FieldMetadata {
pub name: String,
pub data_type: String,
pub nullable: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_length: Option<i32>,
}
impl<T: Serialize> ApiResponse<T> {
pub fn success(data: T, metadata: ResponseMetadata) -> Self {
Self {
success: true,
data: Some(data),
error: None,
metadata,
}
}
pub fn error(code: &str, message: String, metadata: ResponseMetadata) -> Self {
Self {
success: false,
data: None,
error: Some(ErrorInfo {
code: code.to_string(),
message,
details: None,
}),
metadata,
}
}
#[expect(
dead_code,
reason = "legacy response helper retained for API compatibility"
)]
pub fn error_with_details(
code: &str,
message: String,
details: Value,
metadata: ResponseMetadata,
) -> Self {
Self {
success: false,
data: None,
error: Some(ErrorInfo {
code: code.to_string(),
message,
details: Some(details),
}),
metadata,
}
}
}
async fn load_instances() -> anyhow::Result<HashMap<String, PostgresInstance>> {
let mut instances = HashMap::new();
let superuser_password = std::env::var("PG_SUPERUSER_PASSWORD").unwrap_or_default();
instances.insert(
"default".to_string(),
PostgresInstance {
id: "default".to_string(),
name: "Primary Instance".to_string(),
host: std::env::var("PG_HOST").unwrap_or_else(|_| "127.0.0.1".to_string()),
port: std::env::var("PG_PORT")
.ok()
.and_then(|p| p.parse().ok())
.unwrap_or(5432),
superuser: std::env::var("PG_SUPERUSER").unwrap_or_else(|_| "postgres".to_string()),
superuser_password,
instance_type: InstanceType::Single,
created_at: Utc::now(),
status: InstanceStatus::Active,
},
);
Ok(instances)
}
async fn load_accounts() -> anyhow::Result<HashMap<String, Account>> {
let mut accounts = HashMap::new();
let config_dir = std::env::var("CONFIG_DIR").unwrap_or_else(|_| "config".to_string());
let config_path = std::path::PathBuf::from(&config_dir).join("accounts.json");
eprintln!(
"[ACCOUNT LOADING] Loading accounts from: {}",
config_path.display()
);
tracing::info!("Loading accounts from: {}", config_path.display());
match tokio::fs::read_to_string(&config_path).await {
Ok(content) => {
match serde_json::from_str::<Vec<Account>>(&content) {
Ok(loaded_accounts) => {
tracing::info!("Loaded {} accounts from config file", loaded_accounts.len());
for mut account in loaded_accounts {
let (normalized, migrated) = normalize_stored_key(&account.api_key);
if migrated {
tracing::warn!(
"Account '{}' uses a plaintext api_key; \
rotate it and re-save to store only the hash",
account.name
);
}
account.api_key = normalized.clone();
accounts.insert(normalized, account);
}
}
Err(e) => {
tracing::error!("Failed to parse accounts.json: {}", e);
return Err(anyhow::anyhow!(
"Failed to parse accounts configuration: {}",
e
));
}
}
}
Err(e) => {
tracing::error!(
"Could not read accounts file from {}: {}. \
Refusing to start without explicit account configuration \
(run `pg-api setup`).",
config_path.display(),
e
);
return Err(anyhow::anyhow!(
"Missing accounts configuration at {}: {}. \
Create it with `pg-api setup`; the service will not start \
with default credentials.",
config_path.display(),
e
));
}
}
if accounts.is_empty() {
return Err(anyhow::anyhow!(
"No accounts configured in {}. Refusing to start with zero accounts.",
config_path.display()
));
}
Ok(accounts)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_api_response_success() {
let metadata = ResponseMetadata {
request_id: "test-123".to_string(),
execution_time_ms: 42,
rows_affected: Some(1),
instance_id: Some("acc_001".to_string()),
timestamp: Utc::now(),
};
let data = json!({"message": "success"});
let response = ApiResponse::success(data, metadata);
assert!(response.success);
assert!(response.data.is_some());
assert!(response.error.is_none());
assert_eq!(response.metadata.execution_time_ms, 42);
}
#[test]
fn test_api_response_error() {
let metadata = ResponseMetadata {
request_id: "test-456".to_string(),
execution_time_ms: 10,
rows_affected: None,
instance_id: Some("acc_001".to_string()),
timestamp: Utc::now(),
};
let response: ApiResponse<Value> =
ApiResponse::error("PERMISSION_DENIED", "Access denied".to_string(), metadata);
assert!(!response.success);
assert!(response.data.is_none());
assert!(response.error.is_some());
let error = response.error.unwrap();
assert_eq!(error.code, "PERMISSION_DENIED");
assert_eq!(error.message, "Access denied");
}
#[test]
fn test_query_request_deserialization() {
let json_str = r#"{
"query": "SELECT * FROM users WHERE id = $1",
"database": "mydb",
"params": [1]
}"#;
let request: QueryRequest = serde_json::from_str(json_str).unwrap();
assert_eq!(request.query, "SELECT * FROM users WHERE id = $1");
assert_eq!(request.database, "mydb");
assert_eq!(request.params.len(), 1);
assert_eq!(request.params[0], 1);
}
#[test]
fn test_query_request_with_options() {
let json_str = r#"{
"query": "SELECT * FROM users",
"database": "mydb",
"params": [],
"options": {
"timeout_ms": 5000,
"read_only": true
}
}"#;
let request: QueryRequest = serde_json::from_str(json_str).unwrap();
assert_eq!(request.options.timeout_ms, Some(5000));
assert!(request.options.read_only);
}
#[test]
fn test_permission_serialization() {
let perms = vec![Permission::Select, Permission::Insert];
let json = serde_json::to_string(&perms).unwrap();
assert_eq!(json, "[\"SELECT\",\"INSERT\"]");
}
#[test]
fn test_database_management_permission_serialization() {
let perms = vec![Permission::CreateDatabase, Permission::DropOwnedDatabase];
let json = serde_json::to_string(&perms).unwrap();
assert_eq!(json, "[\"CREATE_DATABASE\",\"DROP_OWNED_DATABASE\"]");
}
#[test]
fn test_account_role_deserialization() {
let json_str = "\"superuser\"";
let role: AccountRole = serde_json::from_str(json_str).unwrap();
assert_eq!(role, AccountRole::Superuser);
let role: AccountRole = serde_json::from_str("\"owner\"").unwrap();
assert_eq!(role, AccountRole::Superuser);
let role: AccountRole = serde_json::from_str("\"read_only\"").unwrap();
assert_eq!(role, AccountRole::Readonly);
let role: AccountRole = serde_json::from_str("\"app\"").unwrap();
assert_eq!(role, AccountRole::Application);
let json_str = "\"admin\"";
let role: AccountRole = serde_json::from_str(json_str).unwrap();
assert_eq!(role, AccountRole::Admin);
let json_str = "\"developer\"";
let role: AccountRole = serde_json::from_str(json_str).unwrap();
assert_eq!(role, AccountRole::Developer);
let json_str = "\"application\"";
let role: AccountRole = serde_json::from_str(json_str).unwrap();
assert_eq!(role, AccountRole::Application);
let json_str = "\"readonly\"";
let role: AccountRole = serde_json::from_str(json_str).unwrap();
assert_eq!(role, AccountRole::Readonly);
}
#[test]
fn test_database_access_serialization() {
let access = DatabaseAccess {
database: "testdb".to_string(),
username: "testuser".to_string(),
password: "testpass".to_string(),
permissions: vec![Permission::Select, Permission::All],
};
let json = serde_json::to_string(&access).unwrap();
assert!(json.contains("testdb"));
assert!(json.contains("SELECT"));
assert!(json.contains("ALL"));
}
#[test]
fn test_instance_type_variants() {
let single = InstanceType::Single;
let primary = InstanceType::Primary;
let replica = InstanceType::Replica;
assert!(matches!(single, InstanceType::Single));
assert!(matches!(primary, InstanceType::Primary));
assert!(matches!(replica, InstanceType::Replica));
}
#[test]
fn test_instance_status_variants() {
assert!(matches!(InstanceStatus::Active, InstanceStatus::Active));
assert!(matches!(
InstanceStatus::Maintenance,
InstanceStatus::Maintenance
));
assert!(matches!(InstanceStatus::Degraded, InstanceStatus::Degraded));
assert!(matches!(InstanceStatus::Offline, InstanceStatus::Offline));
}
}