use crate::prelude::*;
use convert_case::{Case, Casing};
use crc32fast::hash as crc32;
use std::fmt;
use tokio_postgres::types::Type as dbtype;
#[derive(Clone)]
pub struct KlaczDB {
pub handle: &'static str,
}
impl KlaczDB {
pub const GET_INSTANCE_ID: &str = "SELECT nextval('_instance_id')";
async fn get_instance_id(&self) -> anyhow::Result<i64> {
let client = DBPools::get_client(self.handle).await?;
let statement = client.prepare_cached(Self::GET_INSTANCE_ID).await?;
let row = client.query_one(&statement, &[]).await?;
row.try_get(0)
.map_err(|e: tokio_postgres::Error| anyhow!(e))
}
pub const GET_LEVEL: &str = r"SELECT _level::bigint
FROM _level
WHERE
_channel = $1
AND _account = $2
LIMIT 1";
pub async fn get_level(&self, room: &Room, user: &UserId) -> anyhow::Result<i64> {
let name = room_name(room);
let user_name = user.as_str();
let client = DBPools::get_client(self.handle).await?;
let statement = client
.prepare_typed_cached(Self::GET_LEVEL, &[dbtype::VARCHAR, dbtype::VARCHAR])
.await?;
let Ok(row) = client.query_one(&statement, &[&name, &user_name]).await else {
return Ok(0);
};
row.try_get(0)
.map_err(|e: tokio_postgres::Error| anyhow!(e))
}
pub const DELETE_LEVELS: &str = "DELETE FROM _level WHERE _channel = $1 and _account = $2";
pub const INSERT_LEVELS: &str = r"INSERT INTO _level (_oid, _channel, _account, _level)
VALUES ($1, $2, $3, $4)";
pub async fn add_level(&self, room: &Room, user: &UserId, level: i64) -> anyhow::Result<()> {
let name = room_name(room);
let user_name = user.as_str();
let id = self.get_instance_id().await?;
let oid = KlaczClass::Level.make_oid(id);
let mut client = DBPools::get_client(self.handle).await?;
let transaction = client.transaction().await?;
let delete = transaction
.prepare_typed_cached(Self::DELETE_LEVELS, &[dbtype::VARCHAR, dbtype::VARCHAR])
.await?;
let insert = transaction
.prepare_typed_cached(
Self::INSERT_LEVELS,
&[dbtype::INT8, dbtype::VARCHAR, dbtype::VARCHAR, dbtype::INT8],
)
.await?;
if transaction.execute(&delete, &[&name, &user_name]).await? > 1 {
transaction.rollback().await?;
bail!("too many deleted levels")
};
if transaction
.execute(&insert, &[&oid, &name, &user_name, &level])
.await?
!= 1
{
transaction.rollback().await?;
bail!("too many inserted levels")
};
transaction
.commit()
.await
.map_err(|e: tokio_postgres::Error| anyhow!(e))
}
pub const GET_TERM_OID: &str = r"SELECT _oid
FROM _term
WHERE _name = $1";
pub const GET_TERM_ENTRY: &str = r"SELECT _text
FROM _entry
WHERE _term_oid = $1
ORDER BY random()
LIMIT 1";
pub async fn get_entry(&self, term: &str) -> anyhow::Result<String> {
let client = DBPools::get_client(self.handle).await?;
let term_statement = client
.prepare_typed_cached(Self::GET_TERM_OID, &[dbtype::VARCHAR])
.await?;
let entry_statement = client
.prepare_typed_cached(Self::GET_TERM_ENTRY, &[dbtype::INT8])
.await?;
let term_rows = client.query(&term_statement, &[&term]).await?;
let term_oid: i64 = match term_rows.len() {
0 => return Err(KlaczError::EntryNotFound.into()),
2.. => return Err(KlaczError::DBInconsistency.into()),
1 => term_rows
.first()
.ok_or_else(|| anyhow!("no row returned despite len() == 1"))?
.try_get(0)?,
};
let entry_rows = client.query(&entry_statement, &[&term_oid]).await?;
match entry_rows.len() {
1 => Ok(entry_rows
.first()
.ok_or_else(|| anyhow!("no row returned despite len() == 1"))?
.try_get(0)?),
_ => Err(KlaczError::DBInconsistency.into()),
}
}
pub const REMOVE_TERM: &str = r"DELETE FROM _term WHERE _oid = $1";
pub const REMOVE_ENTRY: &str = r"DELETE FROM _entry
WHERE _oid = (
SELECT max(_oid)
FROM _entry
WHERE
_term_oid = $1
AND _text = $2
)";
pub const COUNT_ENTRIES: &str = r"SELECT count(*) FROM _entry WHERE _term_oid = $1";
pub async fn remove_entry(&self, term: &str, entry: &str) -> anyhow::Result<KlaczKBChange> {
let mut response = KlaczKBChange::Unchanged;
let mut client = DBPools::get_client(self.handle).await?;
let get_term_statement = client
.prepare_typed_cached(Self::GET_TERM_OID, &[dbtype::VARCHAR])
.await?;
let remove_term_statement = client
.prepare_typed_cached(Self::REMOVE_TERM, &[dbtype::INT8])
.await?;
let remove_entry_statement = client
.prepare_typed_cached(Self::REMOVE_ENTRY, &[dbtype::INT8, dbtype::VARCHAR])
.await?;
let count_entries_statement = client
.prepare_typed_cached(Self::COUNT_ENTRIES, &[dbtype::INT8])
.await?;
let transaction = client.transaction().await?;
let term_rows = transaction.query(&get_term_statement, &[&term]).await?;
let term_oid: i64 = match term_rows.len() {
0 => return Err(KlaczError::TermNotFound.into()),
2.. => return Err(KlaczError::DBInconsistency.into()),
1 => term_rows
.first()
.ok_or_else(|| anyhow!("no row returned despite len() == 1"))?
.try_get(0)?,
};
let deleted = transaction
.execute(&remove_entry_statement, &[&term_oid, &entry])
.await?;
if deleted == 0 {
return Ok(response);
};
response = KlaczKBChange::RemovedEntry;
let left: i64 = transaction
.query_one(&count_entries_statement, &[&term_oid])
.await?
.try_get(0)?;
if left == 0 {
let deleted = transaction
.execute(&remove_term_statement, &[&term_oid])
.await?;
if deleted > 1 {
transaction.rollback().await?;
return Err(KlaczError::DBInconsistency.into());
};
response = KlaczKBChange::RemovedTerm;
};
transaction.commit().await?;
Ok(response)
}
pub const INSERT_TERM: &str = r"INSERT INTO _term (_oid, _name, _visible)
VALUES ($1, $2, true)";
pub const INSERT_ENTRY: &str = r"INSERT INTO _entry (_oid, _term_oid, _added_by, _text, _added_at, _visible)
VALUES ($1, $2, $3, $4, now(), true)";
pub async fn add_entry(
&self,
user: &UserId,
term: &str,
entry: &str,
) -> anyhow::Result<KlaczKBChange> {
let mut client = DBPools::get_client(self.handle).await?;
let mut ok_result = KlaczKBChange::AddedEntry;
let user_name = user.as_str();
let get_term_statement = client
.prepare_typed_cached(Self::GET_TERM_OID, &[dbtype::VARCHAR])
.await?;
let insert_term_statement = client
.prepare_typed_cached(Self::INSERT_TERM, &[dbtype::INT8, dbtype::VARCHAR])
.await?;
let insert_entry_statement = client
.prepare_typed_cached(
Self::INSERT_ENTRY,
&[dbtype::INT8, dbtype::INT8, dbtype::VARCHAR, dbtype::VARCHAR],
)
.await?;
let transaction = client.transaction().await?;
let term_rows = transaction.query(&get_term_statement, &[&term]).await?;
let term_oid: i64 = match term_rows.len() {
2.. => return Err(KlaczError::DBInconsistency.into()),
1 => term_rows
.first()
.ok_or_else(|| anyhow!("no row returned despite len() == 1"))?
.try_get(0)?,
0 => {
let term_instance_id = self.get_instance_id().await?;
let term_oid_new = KlaczClass::Term.make_oid(term_instance_id);
trace!("term: instance_id: {term_instance_id}, oid: {term_oid_new}");
transaction
.execute(&insert_term_statement, &[&term_oid_new, &term])
.await?;
ok_result = KlaczKBChange::CreatedTerm;
term_oid_new
}
};
let entry_instance_id = self.get_instance_id().await?;
let entry_oid = KlaczClass::Entry.make_oid(entry_instance_id);
trace!("entry: instance_id: {entry_instance_id}, oid: {entry_oid}");
transaction
.execute(
&insert_entry_statement,
&[&entry_oid, &term_oid, &user_name, &entry],
)
.await?;
transaction.commit().await?;
Ok(ok_result)
}
}
#[allow(dead_code)]
pub const OID_MAXIMUM_INSTANCE_ID: i64 = 281_474_976_710_655;
pub const OID_MAXIMUM_CLASS_ID: u32 = 65535;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum KlaczKBChange {
CreatedTerm,
AddedEntry,
Unchanged,
RemovedEntry,
RemovedTerm,
}
impl fmt::Display for KlaczKBChange {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{self:?}")
}
}
#[derive(Debug)]
pub enum KlaczClass {
TopicChange,
Term,
Entry,
Level,
Link,
Memo,
Seen,
}
impl fmt::Display for KlaczClass {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{self:?}")
}
}
impl FromStr for KlaczClass {
type Err = KlaczError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let binding = s.to_string().to_case(Case::Kebab);
let unified = binding.as_str();
trace!("binding: {binding}");
trace!("unified: {unified}");
match unified {
"topic-change" => Ok(Self::TopicChange),
"term" => Ok(Self::Term),
"entry" => Ok(Self::Entry),
"level" => Ok(Self::Level),
"link" => Ok(Self::Link),
"memo" => Ok(Self::Memo),
"seen" => Ok(Self::Seen),
_ => Err(KlaczError::UnknownClass),
}
}
}
impl KlaczClass {
#[must_use]
pub fn class_name(&self) -> String {
self.to_string().to_case(Case::UpperKebab)
}
#[must_use]
pub fn class_id(&self) -> i64 {
(crc32(self.class_name().as_bytes()) % OID_MAXIMUM_CLASS_ID).into()
}
pub fn make_oid(&self, instance_id: i64) -> i64 {
let class_id: i64 = self.class_id();
let shifted: i64 = instance_id << 16;
let oid = shifted | class_id;
trace!("class_id: {class_id}; shifted: {shifted}; instance_id: {instance_id}; oid: {oid}");
oid
}
}
#[derive(Debug, PartialEq, Eq, thiserror::Error)]
pub enum KlaczError {
UnknownClass,
TermNotFound,
EntryNotFound,
DBInconsistency,
}
impl fmt::Display for KlaczError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{self:?}")
}
}
fn default_add_keywords() -> Vec<String> {
vec!["add".s()]
}
fn default_remove_keywords() -> Vec<String> {
vec!["remove".s()]
}
#[derive(Clone, Deserialize)]
pub struct ModuleConfig {
#[serde(default = "default_add_keywords")]
pub keywords_add: Vec<String>,
#[serde(default = "default_remove_keywords")]
pub keywords_remove: Vec<String>,
}
pub(crate) fn starter(_: &Client, config: &Config) -> anyhow::Result<Vec<ModuleInfo>> {
info!("registering modules");
let module_config: ModuleConfig = config.typed_module_config(module_path!())?;
let (addtx, addrx) = mpsc::channel::<ConsumerEvent>(1);
let add = ModuleInfo {
name: "add".s(),
help: "add an entry to knowledge base".s(),
acl: vec![],
trigger: TriggerType::Keyword(module_config.keywords_add.clone()),
channel: addtx,
error_prefix: Some("error adding entry".s()),
};
add.spawn(addrx, module_config.clone(), add_processor);
let (removetx, removerx) = mpsc::channel::<ConsumerEvent>(1);
let remove = ModuleInfo {
name: "remove".s(),
help: "remove an entry to knowledge base".s(),
acl: vec![Acl::KlaczLevel(10)],
trigger: TriggerType::Keyword(module_config.keywords_remove.clone()),
channel: removetx,
error_prefix: Some("error removing entry".s()),
};
remove.spawn(removerx, module_config, remove_processor);
Ok(vec![add, remove])
}
pub async fn add_processor(event: ConsumerEvent, _: ModuleConfig) -> anyhow::Result<()> {
let Some(body) = event.args else {
event
.room
.send(RoomMessageEventContent::text_plain(
"missing arguments: term, definition",
))
.await?;
bail!("missing arguments")
};
let mut args = body.splitn(2, [' ', '\n']);
let term = args.next().ok_or_else(|| anyhow!("missing arguments"))?;
let Some(definition) = args.next() else {
event
.room
.send(RoomMessageEventContent::text_plain(
"missing arguments: definition",
))
.await?;
bail!("missing arguments")
};
trace!("attempting to add: term: {term}: definition: {definition}");
let mut response = String::new();
let result = event
.klacz
.add_entry(&event.sender, term, definition)
.await?;
if result == KlaczKBChange::CreatedTerm {
response.push_str(format!("Created term \"{term}\"\n").as_str());
};
response.push_str(format!(r#"Added one entry to term "{term}""#).as_str());
event
.room
.send(RoomMessageEventContent::text_plain(response))
.await?;
Ok(())
}
pub async fn remove_processor(event: ConsumerEvent, _: ModuleConfig) -> anyhow::Result<()> {
let Some(body) = event.args else {
event
.room
.send(RoomMessageEventContent::text_plain(
"missing arguments: term, definition",
))
.await?;
bail!("missing arguments")
};
let mut args = body.splitn(2, [' ', '\n']);
let term = args.next().ok_or_else(|| anyhow!("missing arguments"))?;
let Some(definition) = args.next() else {
event
.room
.send(RoomMessageEventContent::text_plain(
"missing arguments: definition",
))
.await?;
bail!("missing arguments")
};
trace!("attempting to remove: term: {term}: definition: {definition}");
let response = event.klacz.remove_entry(term, definition).await?;
let message = match response {
KlaczKBChange::Unchanged => format!("entry not found in {term}"),
KlaczKBChange::RemovedEntry => format!("removed entry from {term}"),
KlaczKBChange::RemovedTerm => format!("last entry, removed {term}"),
_ => format!("unexpected response from klacz, no error: {response}"),
};
event
.room
.send(RoomMessageEventContent::text_plain(message))
.await?;
Ok(())
}