use anyhow::Result;
use sqlx::{AnyPool, Row};
use uuid::Uuid;
pub struct Store {
pool: AnyPool,
}
impl Store {
pub async fn connect(url: &str) -> Result<Self> {
sqlx::any::install_default_drivers();
if url.starts_with("sqlite:") {
Self::ensure_sqlite_file(url)?;
}
let pool = AnyPool::connect(url).await?;
let store = Self { pool };
store.migrate().await?;
Ok(store)
}
fn ensure_sqlite_file(url: &str) -> Result<()> {
let path_part = url
.strip_prefix("sqlite:///")
.or_else(|| url.strip_prefix("sqlite://"))
.or_else(|| url.strip_prefix("sqlite:"))
.unwrap_or("");
let path_str = path_part.split('?').next().unwrap_or("").trim();
if path_str.is_empty() || path_str == ":memory:" {
return Ok(());
}
let path = std::path::Path::new(path_str);
if let Some(parent) = path.parent() {
if !parent.as_os_str().is_empty() {
std::fs::create_dir_all(parent)?;
}
}
if !path.exists() {
std::fs::File::create(path)?;
}
Ok(())
}
async fn migrate(&self) -> Result<()> {
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS clients (
vtoken TEXT PRIMARY KEY,
name TEXT NOT NULL UNIQUE,
label TEXT,
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
last_seen TIMESTAMPTZ
)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS routing_state (
from_user TEXT PRIMARY KEY,
active_vtoken TEXT NOT NULL,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS context_token_map (
vctx TEXT PRIMARY KEY,
real_ctx TEXT NOT NULL,
expires_at TIMESTAMPTZ
)
"#,
)
.execute(&self.pool)
.await?;
let _ = sqlx::query(
"ALTER TABLE context_token_map ADD COLUMN peer_user_id TEXT NOT NULL DEFAULT ''",
)
.execute(&self.pool)
.await;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS bot_credentials (
id INTEGER PRIMARY KEY,
token TEXT NOT NULL,
base_url TEXT NOT NULL DEFAULT 'https://ilinkai.weixin.qq.com',
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
)
"#,
)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn upsert_client(&self, vtoken: &str, name: &str, label: Option<&str>) -> Result<()> {
sqlx::query(
r#"
INSERT INTO clients (vtoken, name, label)
VALUES ($1, $2, $3)
ON CONFLICT (name) DO UPDATE
SET label = EXCLUDED.label,
last_seen = CURRENT_TIMESTAMP
"#,
)
.bind(vtoken)
.bind(name)
.bind(label)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn touch_client(&self, vtoken: &str) -> Result<()> {
sqlx::query("UPDATE clients SET last_seen = CURRENT_TIMESTAMP WHERE vtoken = $1")
.bind(vtoken)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn list_clients(&self) -> Result<Vec<ClientRow>> {
let rows = sqlx::query("SELECT vtoken, name, label, last_seen FROM clients ORDER BY name")
.fetch_all(&self.pool)
.await?;
Ok(rows
.into_iter()
.map(|r| ClientRow {
vtoken: r.get("vtoken"),
name: r.get("name"),
label: r.get("label"),
last_seen: r.get::<Option<String>, _>("last_seen"),
})
.collect())
}
pub async fn get_client_by_name(&self, name: &str) -> Result<Option<ClientRow>> {
let row = sqlx::query("SELECT vtoken, name, label, last_seen FROM clients WHERE name = $1")
.bind(name)
.fetch_optional(&self.pool)
.await?;
Ok(row.map(|r| ClientRow {
vtoken: r.get("vtoken"),
name: r.get("name"),
label: r.get("label"),
last_seen: r.get::<Option<String>, _>("last_seen"),
}))
}
pub async fn list_routes(&self) -> Result<Vec<(String, String)>> {
let rows = sqlx::query("SELECT from_user, active_vtoken FROM routing_state")
.fetch_all(&self.pool)
.await?;
Ok(rows
.into_iter()
.map(|r| {
(
r.get::<String, _>("from_user"),
r.get::<String, _>("active_vtoken"),
)
})
.collect())
}
pub async fn get_route(&self, from_user: &str) -> Result<Option<String>> {
let row = sqlx::query("SELECT active_vtoken FROM routing_state WHERE from_user = $1")
.bind(from_user)
.fetch_optional(&self.pool)
.await?;
Ok(row.map(|r| r.get("active_vtoken")))
}
pub async fn set_route(&self, from_user: &str, vtoken: &str) -> Result<()> {
sqlx::query(
r#"
INSERT INTO routing_state (from_user, active_vtoken)
VALUES ($1, $2)
ON CONFLICT (from_user) DO UPDATE
SET active_vtoken = EXCLUDED.active_vtoken,
updated_at = CURRENT_TIMESTAMP
"#,
)
.bind(from_user)
.bind(vtoken)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn map_context_token(&self, real_ctx: &str, peer_user_id: &str) -> Result<String> {
let existing = sqlx::query("SELECT vctx FROM context_token_map WHERE real_ctx = $1")
.bind(real_ctx)
.fetch_optional(&self.pool)
.await?;
if let Some(row) = existing {
return Ok(row.get("vctx"));
}
let vctx = format!("vctx_{}", Uuid::new_v4().simple());
sqlx::query(
"INSERT INTO context_token_map (vctx, real_ctx, peer_user_id) VALUES ($1, $2, $3)",
)
.bind(&vctx)
.bind(real_ctx)
.bind(peer_user_id)
.execute(&self.pool)
.await?;
Ok(vctx)
}
pub async fn persist_context_token(
&self,
vctx: &str,
real_ctx: &str,
peer_user_id: &str,
) -> Result<()> {
sqlx::query(
r#"
INSERT INTO context_token_map (vctx, real_ctx, peer_user_id)
VALUES ($1, $2, $3)
ON CONFLICT (vctx) DO NOTHING
"#,
)
.bind(vctx)
.bind(real_ctx)
.bind(peer_user_id)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn resolve_context_token(&self, vctx: &str) -> Result<Option<String>> {
let row = sqlx::query("SELECT real_ctx FROM context_token_map WHERE vctx = $1")
.bind(vctx)
.fetch_optional(&self.pool)
.await?;
Ok(row.map(|r| r.get("real_ctx")))
}
pub async fn resolve_context_token_full(&self, vctx: &str) -> Result<Option<(String, String)>> {
let row = sqlx::query(
"SELECT real_ctx, COALESCE(peer_user_id, '') AS peer_user_id \
FROM context_token_map WHERE vctx = $1",
)
.bind(vctx)
.fetch_optional(&self.pool)
.await?;
Ok(row.map(|r| (r.get("real_ctx"), r.get("peer_user_id"))))
}
pub async fn list_recent_context_tokens(
&self,
limit: i64,
) -> Result<Vec<(String, String, String)>> {
let rows = sqlx::query(
"SELECT vctx, real_ctx, COALESCE(peer_user_id, '') AS peer_user_id \
FROM context_token_map LIMIT $1",
)
.bind(limit)
.fetch_all(&self.pool)
.await?;
Ok(rows
.into_iter()
.map(|r| {
(
r.get::<String, _>("vctx"),
r.get::<String, _>("real_ctx"),
r.get::<String, _>("peer_user_id"),
)
})
.collect())
}
pub async fn save_credentials(&self, token: &str, base_url: &str) -> Result<()> {
sqlx::query(
r#"
INSERT INTO bot_credentials (id, token, base_url)
VALUES (1, $1, $2)
ON CONFLICT (id) DO UPDATE
SET token = EXCLUDED.token,
base_url = EXCLUDED.base_url,
updated_at = CURRENT_TIMESTAMP
"#,
)
.bind(token)
.bind(base_url)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn load_credentials(&self) -> Result<Option<(String, String)>> {
let row = sqlx::query("SELECT token, base_url FROM bot_credentials WHERE id = 1")
.fetch_optional(&self.pool)
.await?;
Ok(row.map(|r| (r.get("token"), r.get("base_url"))))
}
}
#[derive(Debug, Clone)]
pub struct ClientRow {
pub vtoken: String,
pub name: String,
pub label: Option<String>,
pub last_seen: Option<String>,
}