use serde_json::{json, Value};
use crate::error::ImError;
pub(crate) fn require_str_array(args: &Value, key: &str, cmd: &str) -> Result<Value, ImError> {
args.get(key)
.and_then(Value::as_array)
.filter(|a| !a.is_empty() && a.iter().all(Value::is_string))
.cloned()
.map(Value::Array)
.ok_or_else(|| ImError::Parse(format!("{cmd}: 缺/空 {key}(非空字符串数组)")))
}
const MAX_SEARCH_PAGE_SIZE: u64 = 50;
const MAX_SEARCH_CURSOR_BYTES: usize = 4096;
const MAX_SEARCH_REQ_ID_BYTES: usize = 128;
const MAX_SEARCH_TIMESTAMP: u64 = i64::MAX as u64;
const MAX_SEARCH_KEYWORD_CHARS: usize = 512;
const MAX_SEARCH_FILTER_VALUES: usize = 100;
const MAX_SEARCH_IDENTIFIER_BYTES: usize = 128;
fn validate_search_meta(args: &Value, cmd: &str) -> Result<(), ImError> {
args.get("req_id")
.and_then(Value::as_str)
.filter(|value| !value.trim().is_empty() && value.len() <= MAX_SEARCH_REQ_ID_BYTES)
.ok_or_else(|| {
ImError::Parse(format!(
"{cmd}: req_id 必须为 1..={MAX_SEARCH_REQ_ID_BYTES} 字节的非空字符串"
))
})?;
for field in ["user_id", "team_id", "company_id", "sort"] {
if args.get(field).is_some() {
return Err(ImError::Parse(format!(
"{cmd}: {field} 由服务端权威决定,禁止前端下发"
)));
}
}
Ok(())
}
fn append_search_page(args: &Value, cmd: &str, body: &mut Value) -> Result<(), ImError> {
if let Some(cursor) = args.get("cursor") {
let cursor = cursor
.as_str()
.filter(|value| !value.is_empty() && value.len() <= MAX_SEARCH_CURSOR_BYTES)
.ok_or_else(|| {
ImError::Parse(format!(
"{cmd}: cursor 必须为 1..={MAX_SEARCH_CURSOR_BYTES} 字节的字符串"
))
})?;
body["cursor"] = json!(cursor);
}
if let Some(page_size) = args.get("page_size") {
let page_size = page_size
.as_u64()
.filter(|value| (1..=MAX_SEARCH_PAGE_SIZE).contains(value))
.ok_or_else(|| {
ImError::Parse(format!(
"{cmd}: page_size 必须为 1..={MAX_SEARCH_PAGE_SIZE} 的整数"
))
})?;
body["pageSize"] = json!(page_size);
}
Ok(())
}
fn optional_search_str_array(args: &Value, key: &str, cmd: &str) -> Result<Option<Value>, ImError> {
let Some(value) = args.get(key) else {
return Ok(None);
};
value
.as_array()
.filter(|items| {
!items.is_empty()
&& items.len() <= MAX_SEARCH_FILTER_VALUES
&& items
.iter()
.all(|item| {
item.as_str().is_some_and(|text| {
let text = text.trim();
!text.is_empty() && text.len() <= MAX_SEARCH_IDENTIFIER_BYTES
})
})
})
.cloned()
.map(Value::Array)
.map(Some)
.ok_or_else(|| {
ImError::Parse(format!(
"{cmd}: {key} 必须为 1..={MAX_SEARCH_FILTER_VALUES} 个非空字符串,单值不超过 {MAX_SEARCH_IDENTIFIER_BYTES} 字节"
))
})
}
fn validate_search_keyword(keyword: &str, cmd: &str) -> Result<(), ImError> {
if keyword.chars().count() > MAX_SEARCH_KEYWORD_CHARS {
return Err(ImError::Parse(format!(
"{cmd}: keyword 不得超过 {MAX_SEARCH_KEYWORD_CHARS} 个字符"
)));
}
Ok(())
}
fn optional_search_timestamp(args: &Value, key: &str, cmd: &str) -> Result<Option<u64>, ImError> {
let Some(value) = args.get(key) else {
return Ok(None);
};
value
.as_u64()
.filter(|value| *value <= MAX_SEARCH_TIMESTAMP)
.map(Some)
.ok_or_else(|| {
ImError::Parse(format!(
"{cmd}: {key} 必须为 0..={MAX_SEARCH_TIMESTAMP} 的整数"
))
})
}
fn post_search_body(args: &Value, cmd: &str) -> Result<Value, ImError> {
if let Some(scope) = args.get("scope") {
scope
.as_str()
.filter(|scope| matches!(*scope, "channel" | "global"))
.ok_or_else(|| ImError::Parse(format!("{cmd}: scope 必须为 channel 或 global")))?;
}
let channel_id = match args.get("channel_id") {
None => None,
Some(value) => Some(
value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty() && value.len() <= MAX_SEARCH_IDENTIFIER_BYTES)
.ok_or_else(|| {
ImError::Parse(format!(
"{cmd}: channel_id 必须为 1..={MAX_SEARCH_IDENTIFIER_BYTES} 字节的字符串"
))
})?,
),
};
let keyword = match args.get("keyword") {
None => None,
Some(value) => Some(
value
.as_str()
.ok_or_else(|| ImError::Parse(format!("{cmd}: keyword 必须为字符串")))?
.trim(),
)
.filter(|value| !value.is_empty()),
};
if let Some(keyword) = keyword {
validate_search_keyword(keyword, cmd)?;
}
let sender_ids = optional_search_str_array(args, "sender_ids", cmd)?;
let content_types = optional_search_str_array(args, "content_types", cmd)?;
let start_at = optional_search_timestamp(args, "start_at", cmd)?;
let end_at = optional_search_timestamp(args, "end_at", cmd)?;
if start_at.zip(end_at).is_some_and(|(start, end)| start > end) {
return Err(ImError::Parse(format!("{cmd}: start_at 不得晚于 end_at")));
}
if keyword.is_none()
&& channel_id.is_none()
&& sender_ids.is_none()
&& content_types.is_none()
&& start_at.is_none()
&& end_at.is_none()
{
return Err(ImError::Parse(format!(
"{cmd}: keyword 为空时至少提供 sender_ids/content_types/时间筛选之一"
)));
}
let mut body = json!({});
if let Some(value) = keyword {
body["keyword"] = json!(value);
}
if let Some(value) = channel_id {
body["channelId"] = json!(value);
}
if let Some(value) = sender_ids {
body["senderUserIds"] = value;
}
if let Some(value) = content_types {
body["contentTypes"] = value;
}
if let Some(value) = start_at {
body["startAt"] = json!(value);
}
if let Some(value) = end_at {
body["endAt"] = json!(value);
}
append_search_page(args, cmd, &mut body)?;
Ok(body)
}
fn keyword_search_body(args: &Value, cmd: &str) -> Result<Value, ImError> {
let keyword = args
.get("keyword")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| ImError::Parse(format!("{cmd}: keyword 必须为非空字符串")))?;
validate_search_keyword(keyword, cmd)?;
let mut body = json!({ "keyword": keyword });
append_search_page(args, cmd, &mut body)?;
Ok(body)
}
fn global_search_body(args: &Value, cmd: &str) -> Result<Value, ImError> {
if args.get("cursor").is_some() {
return Err(ImError::Parse(format!(
"{cmd}: 聚合搜索禁止 cursor,请使用三段独立 cursor"
)));
}
let keyword = args
.get("keyword")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| ImError::Parse(format!("{cmd}: keyword 必须为非空字符串")))?;
validate_search_keyword(keyword, cmd)?;
let mut body = json!({ "keyword": keyword });
append_search_page(args, cmd, &mut body)?;
for (field, wire) in [
("user_cursor", "userCursor"),
("channel_cursor", "channelCursor"),
("message_group_cursor", "messageGroupCursor"),
] {
if let Some(cursor) = args.get(field) {
let cursor = cursor
.as_str()
.filter(|value| !value.is_empty() && value.len() <= MAX_SEARCH_CURSOR_BYTES)
.ok_or_else(|| {
ImError::Parse(format!(
"{cmd}: {field} 必须为 1..={MAX_SEARCH_CURSOR_BYTES} 字节的字符串"
))
})?;
body[wire] = json!(cursor);
}
}
Ok(body)
}
pub(crate) fn do_search_body(args: &Value, cmd: &str) -> Result<Value, ImError> {
validate_search_meta(args, cmd)?;
match cmd {
"im_search_post" => post_search_body(args, cmd),
"im_search_user" | "im_search_channel" => keyword_search_body(args, cmd),
"im_search_do" => global_search_body(args, cmd),
_ => Err(ImError::Parse(format!("{cmd}: 未知搜索命令"))),
}
}