use helix_core::effect::{Effect, Row, ScopedGetSpec, SqlValue, StorageOp, UpsertSpec};
use helix_core::EffectSink;
use serde_json::{Map, Value};
use crate::error::ImError;
use crate::event::MessageV3Event;
use crate::module::ImModule;
use crate::state::CorrelationContext;
pub const DRAFT_TABLE: &str = "message_draft";
pub const DRAFT_RESULT_EVENT: &str = "im:read:result";
const SAVE_COMMAND: &str = "im_save_draft";
const QUERY_COMMAND: &str = "im_query_draft";
const SAVE_KEYS: &[&str] = &["channel_id", "text", "props", "updated_at", "req_id"];
const QUERY_KEYS: &[&str] = &["channel_id", "req_id"];
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DraftProjection {
pub account_id: String,
pub channel_id: String,
pub text: String,
pub props: Map<String, Value>,
pub updated_at: i64,
}
impl DraftProjection {
pub fn to_json(&self) -> Value {
serde_json::json!({
"accountId": self.account_id,
"channelId": self.channel_id,
"text": self.text,
"props": Value::Object(self.props.clone()),
"updatedAt": self.updated_at,
})
}
pub fn to_row(&self) -> Row {
vec![
(
"account_id".to_string(),
SqlValue::Text(self.account_id.clone()),
),
(
"channel_id".to_string(),
SqlValue::Text(self.channel_id.clone()),
),
("text".to_string(), SqlValue::Text(self.text.clone())),
(
"props".to_string(),
SqlValue::Text(
serde_json::to_string(&Value::Object(self.props.clone()))
.unwrap_or_else(|_| "{}".to_string()),
),
),
("updated_at".to_string(), SqlValue::Integer(self.updated_at)),
]
}
pub fn from_row(row: &Row) -> Option<Self> {
let account_id = text_column(row, "account_id")?.to_string();
let channel_id = text_column(row, "channel_id")?.to_string();
if account_id.is_empty() || channel_id.is_empty() {
return None;
}
Some(Self {
account_id,
channel_id,
text: text_column(row, "text").unwrap_or_default().to_string(),
props: text_column(row, "props")
.and_then(|raw| serde_json::from_str::<Value>(raw).ok())
.and_then(|value| value.as_object().cloned())
.unwrap_or_default(),
updated_at: integer_column(row, "updated_at").unwrap_or_default(),
})
}
pub fn upsert_op(&self) -> StorageOp {
StorageOp::BatchUpsert(UpsertSpec {
version_column: None,
update_guard: None,
table: DRAFT_TABLE,
rows: vec![self.to_row()],
conflict_key: None,
exclude_from_update: Vec::new(),
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SaveDraftCommand {
pub draft: DraftProjection,
pub req_id: Option<String>,
}
impl SaveDraftCommand {
pub fn parse(payload: &[u8], account_id: &str, now_ms: u64) -> Result<Self, ImError> {
let obj = as_object(payload, SAVE_COMMAND)?;
reject_unknown(&obj, SAVE_KEYS, SAVE_COMMAND)?;
let account_id = require_account(account_id, SAVE_COMMAND)?;
let channel_id = require_channel_id(&obj, SAVE_COMMAND)?;
let text = match obj.get("text") {
None | Some(Value::Null) => String::new(),
Some(Value::String(text)) => text.clone(),
Some(_) => {
return Err(ImError::Parse(format!("{SAVE_COMMAND}: text 必须是字符串")));
}
};
let props = match obj.get("props") {
None | Some(Value::Null) => Map::new(),
Some(Value::Object(props)) => props.clone(),
Some(_) => {
return Err(ImError::Parse(format!("{SAVE_COMMAND}: props 必须是对象")));
}
};
let updated_at = match obj.get("updated_at") {
None | Some(Value::Null) => now_ms as i64,
Some(value) => value.as_i64().filter(|value| *value > 0).ok_or_else(|| {
ImError::Parse(format!("{SAVE_COMMAND}: updated_at 必须是正整数"))
})?,
};
Ok(Self {
draft: DraftProjection {
account_id,
channel_id,
text,
props,
updated_at,
},
req_id: optional_req_id(&obj),
})
}
pub fn result_event(&self) -> Result<MessageV3Event, ImError> {
result_event(self.req_id.as_deref(), Some(&self.draft))
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct QueryDraftCommand {
pub account_id: String,
pub channel_id: String,
pub req_id: Option<String>,
}
impl QueryDraftCommand {
pub fn parse(payload: &[u8], account_id: &str) -> Result<Self, ImError> {
let obj = as_object(payload, QUERY_COMMAND)?;
reject_unknown(&obj, QUERY_KEYS, QUERY_COMMAND)?;
Ok(Self {
account_id: require_account(account_id, QUERY_COMMAND)?,
channel_id: require_channel_id(&obj, QUERY_COMMAND)?,
req_id: optional_req_id(&obj),
})
}
pub fn get_op(&self) -> StorageOp {
StorageOp::ScopedGet(ScopedGetSpec {
table: DRAFT_TABLE,
scope_col: "account_id",
scope_val: SqlValue::Text(self.account_id.clone()),
key_col: "channel_id",
key_val: SqlValue::Text(self.channel_id.clone()),
})
}
pub fn result_event(&self, rows: &[Row]) -> Result<MessageV3Event, ImError> {
let draft = rows
.iter()
.find_map(DraftProjection::from_row)
.filter(|draft| {
draft.account_id == self.account_id && draft.channel_id == self.channel_id
});
result_event(self.req_id.as_deref(), draft.as_ref())
}
}
fn result_event(
req_id: Option<&str>,
draft: Option<&DraftProjection>,
) -> Result<MessageV3Event, ImError> {
MessageV3Event::new(
DRAFT_RESULT_EVENT,
serde_json::json!({
"req_id": req_id.unwrap_or_default(),
"body": {
"ok": true,
"draft": draft.map(DraftProjection::to_json).unwrap_or(Value::Null),
},
}),
)
}
pub const DRAFT_COMMANDS: &[&str] = &[SAVE_COMMAND, QUERY_COMMAND];
pub fn is_draft_command(name: &str) -> bool {
DRAFT_COMMANDS.contains(&name)
}
pub fn handle_command(
module: &mut ImModule,
name: &str,
payload: &[u8],
now_ms: u64,
out: &mut EffectSink,
) -> Result<(), ImError> {
match name {
SAVE_COMMAND => handle_save_draft(module, payload, now_ms, out),
QUERY_COMMAND => {
handle_query_draft(module, payload, out)?;
Ok(())
}
other => Err(ImError::Parse(format!("draft: 未认领的命令 {other}"))),
}
}
pub fn handle_save_draft(
module: &mut ImModule,
payload: &[u8],
now_ms: u64,
out: &mut EffectSink,
) -> Result<(), ImError> {
let command = SaveDraftCommand::parse(payload, module.config.auth_user_id.as_str(), now_ms)?;
let terminal = command.result_event()?.into_bytes();
let corr = module.alloc_corr_internal();
module.state.corr_map.insert(
corr,
CorrelationContext::MessageV3Commit {
terminal_events: vec![terminal],
},
);
out.push(Effect::Persist {
corr,
ops: vec![command.draft.upsert_op()],
});
Ok(())
}
pub fn handle_query_draft(
module: &mut ImModule,
payload: &[u8],
out: &mut EffectSink,
) -> Result<QueryDraftCommand, ImError> {
let command = QueryDraftCommand::parse(payload, module.config.auth_user_id.as_str())?;
let corr = module.alloc_corr_internal();
module.state.corr_map.insert(
corr,
CorrelationContext::MessageV3DraftReadback {
command: Box::new(command.clone()),
},
);
out.push(Effect::Persist {
corr,
ops: vec![command.get_op()],
});
Ok(command)
}
fn as_object(payload: &[u8], cmd: &str) -> Result<Map<String, Value>, ImError> {
let value: Value = serde_json::from_slice(payload)
.map_err(|error| ImError::Parse(format!("{cmd} payload: {error}")))?;
value
.as_object()
.cloned()
.ok_or_else(|| ImError::Parse(format!("{cmd} payload 必须是对象")))
}
fn reject_unknown(obj: &Map<String, Value>, allowed: &[&str], cmd: &str) -> Result<(), ImError> {
for key in obj.keys() {
if !allowed.contains(&key.as_str()) {
return Err(ImError::Parse(format!("{cmd}: 未知字段 {key}")));
}
}
Ok(())
}
fn require_account(account_id: &str, cmd: &str) -> Result<String, ImError> {
if account_id.is_empty() {
return Err(ImError::Parse(format!("{cmd}: 缺运行时账号身份")));
}
Ok(account_id.to_string())
}
fn require_channel_id(obj: &Map<String, Value>, cmd: &str) -> Result<String, ImError> {
obj.get("channel_id")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.map(str::to_string)
.ok_or_else(|| ImError::Parse(format!("{cmd}: 缺/空 channel_id")))
}
fn optional_req_id(obj: &Map<String, Value>) -> Option<String> {
obj.get("req_id")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.map(str::to_string)
}
fn text_column<'a>(row: &'a Row, column: &str) -> Option<&'a str> {
row.iter().find_map(|(name, value)| match value {
SqlValue::Text(value) if name == column => Some(value.as_str()),
_ => None,
})
}
fn integer_column(row: &Row, column: &str) -> Option<i64> {
row.iter().find_map(|(name, value)| match value {
SqlValue::Integer(value) if name == column => Some(*value),
_ => None,
})
}