Compare commits
3
Commits
33db3e7c84
...
7582e8b55e
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7582e8b55e | ||
|
|
f9351e8b61 | ||
|
|
31d92f841d |
+1
-5
@@ -35,16 +35,12 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
|||||||
return fmt.Errorf("mkdir data_dir: %w", err)
|
return fmt.Errorf("mkdir data_dir: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
db, err := store.OpenWriter(cfg.DataDir, cfg.SQLiteSynchronous)
|
db, err := store.Open(cfg.DataDir, cfg.SQLiteSynchronous)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer func() { _ = db.Close() }()
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
if migErr := store.Migrate(db); migErr != nil {
|
|
||||||
return migErr
|
|
||||||
}
|
|
||||||
|
|
||||||
mux := http.NewServeMux()
|
mux := http.NewServeMux()
|
||||||
mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, _ *http.Request) {
|
mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, _ *http.Request) {
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
|
|||||||
@@ -43,6 +43,60 @@
|
|||||||
- 备选方案:docker 目标仅 echo 提示。
|
- 备选方案:docker 目标仅 echo 提示。
|
||||||
- 影响:镜像发布流程仍由 Q4 定稿。
|
- 影响:镜像发布流程仍由 Q4 定稿。
|
||||||
|
|
||||||
|
### T0.2 2026-09-30
|
||||||
|
|
||||||
|
1. **请求指纹规范化格式**
|
||||||
|
- 原条款:DEVELOPMENT 7.3「下列字段规范化后的 SHA-256」,未规定字节布局。
|
||||||
|
- 实际做法:对 `to.kind`、`to.id`、`body.enc`、有效 `content_type`、解码后正文、`meta` 规范 JSON、`send_at_ms`、`delay_ms`、有效 `offline.keep`/`ttl_seconds`、有效 `receipt` 做长度前缀(或有无标记)串联后算 SHA-256,输出小写十六进制;`meta` 用 `encoding/json` 对 `map` 键排序序列化;不含 `talk_password`、`rid`。
|
||||||
|
- 原因:文档未给规范格式,需固定、与键顺序无关、含可选字段区分。
|
||||||
|
- 备选方案:整段规范 JSON 对象再哈希。
|
||||||
|
- 影响:各语言 SDK / 服务端必须共用同一布局,否则防重失效。
|
||||||
|
|
||||||
|
2. **指纹使用文档默认值后的有效字段**
|
||||||
|
- 原条款:指纹字段列表含 `receipt`、`offline.*`,未说明缺省如何表示。
|
||||||
|
- 实际做法:`receipt` 缺省按 true;`offline.keep` 缺省 false;`keep` 为 true 且未给 `ttl_seconds` 时按 86400;`keep` 为 false 时 ttl 记 0;`content_type` 按 enc 补默认。
|
||||||
|
- 原因:重试时省略与显式默认应视为同一请求。
|
||||||
|
- 备选方案:按原始 JSON 有无字段区分,省略与显式默认算冲突。
|
||||||
|
- 影响:SDK 省略默认字段时防重仍命中。
|
||||||
|
|
||||||
|
3. **JSON Encoder 去掉尾部换行**
|
||||||
|
- 原条款:用 `json.Encoder` 且 `SetEscapeHTML(false)`。
|
||||||
|
- 实际做法:Encode 后去掉 `Encoder.Encode` 追加的 `\n`,整帧字节数不含该换行。
|
||||||
|
- 原因:MQTT 一发布一帧,示例 JSON 无尾换行;保留换行会抬高帧长并与本地 `frame_too_large` 判断不一致。
|
||||||
|
- 备选方案:保留换行并在 DEVELOPMENT 写明。
|
||||||
|
- 影响:线上帧比「裸 Encoder.Encode」少 1 字节。
|
||||||
|
|
||||||
|
4. **协议包校验范围**
|
||||||
|
- 原条款:T0.2 要求编号规则、正文/meta 大小、`send_at_ms`/`delay_ms` 互斥、登录密码不以 `nst_` 开头。
|
||||||
|
- 实际做法:上述必做之外,顺带校验各帧 `v`/`type`/`rid`、目标 kind、分页 limit、致命 reason 枚举等结构字段;业务错误(目标不存在、配额等)不在本包判定。
|
||||||
|
- 原因:无结构校验则编解码测试无法覆盖「合法帧」边界。
|
||||||
|
- 备选方案:协议包只做编解码,校验留给各 app 模块。
|
||||||
|
- 影响:服务端应复用本包 `Validate`,避免重复规则。
|
||||||
|
|
||||||
|
### T0.3 2026-09-30
|
||||||
|
|
||||||
|
1. **完整表放在 0002,不改已发布的 0001**
|
||||||
|
- 原条款:TASKS T0.3 / 4.2「`0001_init.sql` 包含第 7.7 节全部表」;T0.1 偏差曾写「T0.3 需替换 0001 正文」。
|
||||||
|
- 实际做法:保留 `0001_init.sql` 为 `SELECT 1;`;新增 `0002_schema.sql` 写入 DEVELOPMENT 7.7 全部业务表与索引(含 `api_tokens`、`settings`、`session_hash` 等)。`schema_migrations` 仍由迁移执行器 `CREATE TABLE IF NOT EXISTS` 维护,不放入 0002。
|
||||||
|
- 原因:T0.1 的 0001 可能已记入已有库的 `schema_migrations`;改写已发布迁移语义会导致「版本已应用但表不存在」。
|
||||||
|
- 备选方案:对未迁移库特殊检测并改写 0001(复杂且易错)。
|
||||||
|
- 影响:新库会有版本 1+2 两行;与 TASKS「表在 0001」字面不一致,与「不改已发布迁移」一致。
|
||||||
|
|
||||||
|
2. **写入队列先做一操作一事务**
|
||||||
|
- 原条款:DEVELOPMENT 7.2 合并提交(最多 256 或凑满 2ms,SAVEPOINT);TASKS T0.3 允许简单实现,P2 换合并。
|
||||||
|
- 实际做法:`store.Queue` 用互斥锁串行,每请求一个事务;注释与本条标明 P2 再改为写 goroutine 合并提交。
|
||||||
|
- 原因:本任务范围;合并留给平台 P2。
|
||||||
|
- 备选方案:T0.3 直接做合并(抢 P2 范围)。
|
||||||
|
- 影响:高并发写入落盘次数偏多,正式压测前需完成 P2。
|
||||||
|
|
||||||
|
3. **空库不备份;仅已有 db 文件且有未应用版本时 VACUUM INTO**
|
||||||
|
- 原条款:DEVELOPMENT 7.7「有未应用版本时先 VACUUM INTO」;未区分空库。
|
||||||
|
- 实际做法:`Open` 在打开前检查 `nixmsg.db` 是否已存在;不存在则跳过备份;存在且有 pending 则写入 `<data_dir>/backup/pre-migrate-<UTC时间>.db`。迁移失败返回错误,不自动从备份恢复。
|
||||||
|
- 原因:空库备份无意义;失败退出与文档一致,恢复交给运维。
|
||||||
|
- 备选方案:失败时自动还原备份再退出。
|
||||||
|
- 影响:与任务说明一致;运维需知备份路径。
|
||||||
|
|
||||||
|
|
||||||
## 平台 P
|
## 平台 P
|
||||||
|
|
||||||
暂无。
|
暂无。
|
||||||
|
|||||||
@@ -0,0 +1,94 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
// 协议常量与默认上限(与 DEVELOPMENT 第 6、11 节示例一致)。
|
||||||
|
|
||||||
|
const (
|
||||||
|
Version = 1
|
||||||
|
|
||||||
|
TypeHello = "hello"
|
||||||
|
TypeResp = "resp"
|
||||||
|
TypeSend = "send"
|
||||||
|
TypeMsg = "msg"
|
||||||
|
TypeAck = "ack"
|
||||||
|
TypeRecall = "recall"
|
||||||
|
TypeStatus = "status"
|
||||||
|
TypeReceipt = "receipt"
|
||||||
|
TypeReceiptAck = "receipt_ack"
|
||||||
|
TypeRevoked = "revoked"
|
||||||
|
TypePresenceGet = "presence.get"
|
||||||
|
TypeDirectoryList = "directory.list"
|
||||||
|
TypePresenceWatch = "presence.watch"
|
||||||
|
TypePresence = "presence"
|
||||||
|
TypeUnlock = "unlock"
|
||||||
|
TypeSelfGet = "self.get"
|
||||||
|
TypeSelfUpdate = "self.update"
|
||||||
|
TypeSelfTalkPassword = "self.talk_password"
|
||||||
|
TypeSelfLoginPassword = "self.login_password"
|
||||||
|
TypeSelfLogout = "self.logout"
|
||||||
|
TypeGroupCreate = "group.create"
|
||||||
|
TypeGroupAdd = "group.add"
|
||||||
|
TypeGroupRemove = "group.remove"
|
||||||
|
TypeGroupLeave = "group.leave"
|
||||||
|
TypeGroupTransfer = "group.transfer"
|
||||||
|
TypeGroupRename = "group.rename"
|
||||||
|
TypeGroupDissolve = "group.dissolve"
|
||||||
|
TypeGroupList = "group.list"
|
||||||
|
TypeGroupGet = "group.get"
|
||||||
|
TypeGroupEvent = "group_event"
|
||||||
|
TypeFatal = "fatal"
|
||||||
|
|
||||||
|
TargetEndpoint = "endpoint"
|
||||||
|
TargetGroup = "group"
|
||||||
|
|
||||||
|
EncUTF8 = "utf8"
|
||||||
|
EncBase64 = "base64"
|
||||||
|
|
||||||
|
DefaultContentTypeUTF8 = "text/plain; charset=utf-8"
|
||||||
|
DefaultContentTypeBase64 = "application/octet-stream"
|
||||||
|
|
||||||
|
SessionTokenPrefix = "nst_"
|
||||||
|
|
||||||
|
MinIDLen = 1
|
||||||
|
MaxIDLen = 64
|
||||||
|
MaxContentTypeLen = 128
|
||||||
|
MaxNameChars = 64
|
||||||
|
MinLoginPasswordLen = 8
|
||||||
|
MaxLoginPasswordLen = 128
|
||||||
|
MinTalkPasswordLen = 4
|
||||||
|
MaxTalkPasswordLen = 64
|
||||||
|
MinMaxReceiveBytes = 1024
|
||||||
|
MaxPresenceGetIDs = 200
|
||||||
|
MaxPresenceWatchIDs = 1000
|
||||||
|
MaxPageLimit = 200
|
||||||
|
DefaultOfflineTTL = int64(86400)
|
||||||
|
DefaultMaxBodyBytes = 262144
|
||||||
|
DefaultMaxMetaBytes = 4096
|
||||||
|
DefaultMaxFrameBytes = 786432
|
||||||
|
)
|
||||||
|
|
||||||
|
// 第 6.10 节错误码。
|
||||||
|
const (
|
||||||
|
CodeBadRequest = "bad_request"
|
||||||
|
CodeNotReady = "not_ready"
|
||||||
|
CodeUnauthorized = "unauthorized"
|
||||||
|
CodeForbidden = "forbidden"
|
||||||
|
CodeNotFound = "not_found"
|
||||||
|
CodeInvalidTarget = "invalid_target"
|
||||||
|
CodeConflict = "conflict"
|
||||||
|
CodeIDTaken = "id_taken"
|
||||||
|
CodeBodyTooLarge = "body_too_large"
|
||||||
|
CodeMetaTooLarge = "meta_too_large"
|
||||||
|
CodeFrameTooLarge = "frame_too_large"
|
||||||
|
CodeResponseTooLarge = "response_too_large"
|
||||||
|
CodeTalkPasswordRequired = "talk_password_required"
|
||||||
|
CodeTalkPasswordInvalid = "talk_password_invalid"
|
||||||
|
CodeRateLimited = "rate_limited"
|
||||||
|
CodeNotMember = "not_member"
|
||||||
|
CodeOwnerCannotLeave = "owner_cannot_leave"
|
||||||
|
CodeGroupFull = "group_full"
|
||||||
|
CodeQuotaExceeded = "quota_exceeded"
|
||||||
|
CodeEndpointDisabled = "endpoint_disabled"
|
||||||
|
CodeRegistrationClosed = "registration_closed"
|
||||||
|
CodeRegistrationCodeInvalid = "registration_code_invalid"
|
||||||
|
CodeBusy = "busy"
|
||||||
|
)
|
||||||
@@ -0,0 +1,115 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
type typePeek struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Decode 解析一帧 MQTT/应用 JSON,返回具体类型。
|
||||||
|
func Decode(data []byte) (any, error) {
|
||||||
|
var peek typePeek
|
||||||
|
if err := Unmarshal(data, &peek); err != nil {
|
||||||
|
return nil, badRequest("invalid json")
|
||||||
|
}
|
||||||
|
if peek.V != 0 && peek.V != Version {
|
||||||
|
return nil, badRequest("unsupported version")
|
||||||
|
}
|
||||||
|
if peek.Type == "" {
|
||||||
|
return nil, badRequest("missing type")
|
||||||
|
}
|
||||||
|
|
||||||
|
var out any
|
||||||
|
switch peek.Type {
|
||||||
|
case TypeHello:
|
||||||
|
out = &Hello{}
|
||||||
|
case TypeResp:
|
||||||
|
out = &Resp{}
|
||||||
|
case TypeSend:
|
||||||
|
out = &Send{}
|
||||||
|
case TypeMsg:
|
||||||
|
out = &Msg{}
|
||||||
|
case TypeAck:
|
||||||
|
out = &Ack{}
|
||||||
|
case TypeRecall:
|
||||||
|
out = &Recall{}
|
||||||
|
case TypeStatus:
|
||||||
|
out = &Status{}
|
||||||
|
case TypeReceipt:
|
||||||
|
out = &Receipt{}
|
||||||
|
case TypeReceiptAck:
|
||||||
|
out = &ReceiptAck{}
|
||||||
|
case TypeRevoked:
|
||||||
|
out = &Revoked{}
|
||||||
|
case TypePresenceGet:
|
||||||
|
out = &PresenceGet{}
|
||||||
|
case TypeDirectoryList:
|
||||||
|
out = &DirectoryList{}
|
||||||
|
case TypePresenceWatch:
|
||||||
|
out = &PresenceWatch{}
|
||||||
|
case TypePresence:
|
||||||
|
out = &Presence{}
|
||||||
|
case TypeUnlock:
|
||||||
|
out = &Unlock{}
|
||||||
|
case TypeSelfGet:
|
||||||
|
out = &SelfGet{}
|
||||||
|
case TypeSelfUpdate:
|
||||||
|
out = &SelfUpdate{}
|
||||||
|
case TypeSelfTalkPassword:
|
||||||
|
out = &SelfTalkPassword{}
|
||||||
|
case TypeSelfLoginPassword:
|
||||||
|
out = &SelfLoginPassword{}
|
||||||
|
case TypeSelfLogout:
|
||||||
|
out = &SelfLogout{}
|
||||||
|
case TypeGroupCreate:
|
||||||
|
out = &GroupCreate{}
|
||||||
|
case TypeGroupAdd:
|
||||||
|
out = &GroupAdd{}
|
||||||
|
case TypeGroupRemove:
|
||||||
|
out = &GroupRemove{}
|
||||||
|
case TypeGroupLeave:
|
||||||
|
out = &GroupLeave{}
|
||||||
|
case TypeGroupTransfer:
|
||||||
|
out = &GroupTransfer{}
|
||||||
|
case TypeGroupRename:
|
||||||
|
out = &GroupRename{}
|
||||||
|
case TypeGroupDissolve:
|
||||||
|
out = &GroupDissolve{}
|
||||||
|
case TypeGroupList:
|
||||||
|
out = &GroupList{}
|
||||||
|
case TypeGroupGet:
|
||||||
|
out = &GroupGet{}
|
||||||
|
case TypeGroupEvent:
|
||||||
|
out = &GroupEvent{}
|
||||||
|
case TypeFatal:
|
||||||
|
out = &Fatal{}
|
||||||
|
default:
|
||||||
|
return nil, badRequest(fmt.Sprintf("unknown type %q", peek.Type))
|
||||||
|
}
|
||||||
|
if err := Unmarshal(data, out); err != nil {
|
||||||
|
return nil, badRequest("invalid json for type")
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DecodeRegister 解析注册 HTTP 请求体。
|
||||||
|
func DecodeRegister(data []byte) (*RegisterRequest, error) {
|
||||||
|
var req RegisterRequest
|
||||||
|
if err := Unmarshal(data, &req); err != nil {
|
||||||
|
return nil, badRequest("invalid json")
|
||||||
|
}
|
||||||
|
return &req, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MustRaw 将值编码为 json.RawMessage(用于填 Resp.Data)。
|
||||||
|
func MustRaw(v any) json.RawMessage {
|
||||||
|
b, err := Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Marshal 将值编码为 UTF-8 JSON,不转义 HTML 与非 ASCII,且不含尾部换行。
|
||||||
|
func Marshal(v any) ([]byte, error) {
|
||||||
|
var buf bytes.Buffer
|
||||||
|
if err := Encode(&buf, v); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return buf.Bytes(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Encode 写入 JSON。json.Encoder 默认会追加 '\n',这里去掉,保证整帧字节数与线上一致。
|
||||||
|
func Encode(w io.Writer, v any) error {
|
||||||
|
var buf bytes.Buffer
|
||||||
|
enc := json.NewEncoder(&buf)
|
||||||
|
enc.SetEscapeHTML(false)
|
||||||
|
if err := enc.Encode(v); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
b := buf.Bytes()
|
||||||
|
if n := len(b); n > 0 && b[n-1] == '\n' {
|
||||||
|
b = b[:n-1]
|
||||||
|
}
|
||||||
|
_, err := w.Write(b)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unmarshal 解码 JSON。
|
||||||
|
func Unmarshal(data []byte, v any) error {
|
||||||
|
dec := json.NewDecoder(bytes.NewReader(data))
|
||||||
|
dec.UseNumber()
|
||||||
|
return dec.Decode(v)
|
||||||
|
}
|
||||||
|
|
||||||
|
// FrameBytes 返回整帧编码后的字节数。
|
||||||
|
func FrameBytes(v any) (int, error) {
|
||||||
|
b, err := Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return len(b), nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
// ErrorBody 是 resp / 注册失败里的 error 对象。
|
||||||
|
type ErrorBody struct {
|
||||||
|
Code string `json:"code"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Error 是协议包校验失败时返回的错误,Code 对应第 6.10 节。
|
||||||
|
type Error struct {
|
||||||
|
Code string
|
||||||
|
Message string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *Error) Error() string {
|
||||||
|
if e == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if e.Message == "" {
|
||||||
|
return e.Code
|
||||||
|
}
|
||||||
|
return e.Code + ": " + e.Message
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *Error) Body() ErrorBody {
|
||||||
|
if e == nil {
|
||||||
|
return ErrorBody{}
|
||||||
|
}
|
||||||
|
return ErrorBody{Code: e.Code, Message: e.Message}
|
||||||
|
}
|
||||||
|
|
||||||
|
func badRequest(msg string) *Error {
|
||||||
|
return &Error{Code: CodeBadRequest, Message: msg}
|
||||||
|
}
|
||||||
|
|
||||||
|
func bodyTooLarge(msg string) *Error {
|
||||||
|
return &Error{Code: CodeBodyTooLarge, Message: msg}
|
||||||
|
}
|
||||||
|
|
||||||
|
func metaTooLarge(msg string) *Error {
|
||||||
|
return &Error{Code: CodeMetaTooLarge, Message: msg}
|
||||||
|
}
|
||||||
|
|
||||||
|
func frameTooLarge(msg string) *Error {
|
||||||
|
return &Error{Code: CodeFrameTooLarge, Message: msg}
|
||||||
|
}
|
||||||
@@ -0,0 +1,81 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/binary"
|
||||||
|
"encoding/hex"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RequestFingerprint 计算第 7.3 节发送请求指纹(十六进制小写 SHA-256)。
|
||||||
|
// 字段:to.kind、to.id、body.enc、content_type、解码后正文、meta(键排序)、
|
||||||
|
// send_at_ms、delay_ms、offline.keep、offline.ttl_seconds、receipt。
|
||||||
|
// 不含 talk_password 与 rid。meta 键顺序不影响结果。
|
||||||
|
// receipt / offline 使用文档默认值后的有效值;缺省的 send_at_ms、delay_ms 不写入对应槽位。
|
||||||
|
func RequestFingerprint(s *Send) (string, error) {
|
||||||
|
raw, err := DecodeBody(s.Body)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
metaJSON, err := MetaCanonicalJSON(s.Meta)
|
||||||
|
if err != nil {
|
||||||
|
return "", badRequest("invalid meta")
|
||||||
|
}
|
||||||
|
|
||||||
|
h := sha256.New()
|
||||||
|
writeStr(h, s.To.Kind)
|
||||||
|
writeStr(h, s.To.ID)
|
||||||
|
writeStr(h, s.Body.Enc)
|
||||||
|
writeStr(h, EffectiveContentType(s.Body))
|
||||||
|
writeBytes(h, raw)
|
||||||
|
writeBytes(h, metaJSON)
|
||||||
|
|
||||||
|
writeOptInt64(h, s.SendAtMs)
|
||||||
|
writeOptInt64(h, s.DelayMs)
|
||||||
|
|
||||||
|
keep := EffectiveOfflineKeep(s)
|
||||||
|
writeBool(h, keep)
|
||||||
|
ttl := EffectiveOfflineTTL(s)
|
||||||
|
writeInt64(h, ttl)
|
||||||
|
writeBool(h, EffectiveReceipt(s))
|
||||||
|
|
||||||
|
sum := h.Sum(nil)
|
||||||
|
return hex.EncodeToString(sum), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type hashWriter interface {
|
||||||
|
Write(p []byte) (int, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeStr(h hashWriter, s string) {
|
||||||
|
writeBytes(h, []byte(s))
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeBytes(h hashWriter, b []byte) {
|
||||||
|
var lenBuf [8]byte
|
||||||
|
binary.BigEndian.PutUint64(lenBuf[:], uint64(len(b)))
|
||||||
|
_, _ = h.Write(lenBuf[:])
|
||||||
|
_, _ = h.Write(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeBool(h hashWriter, v bool) {
|
||||||
|
if v {
|
||||||
|
_, _ = h.Write([]byte{1})
|
||||||
|
} else {
|
||||||
|
_, _ = h.Write([]byte{0})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeInt64(h hashWriter, v int64) {
|
||||||
|
var buf [8]byte
|
||||||
|
binary.BigEndian.PutUint64(buf[:], uint64(v))
|
||||||
|
_, _ = h.Write(buf[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeOptInt64(h hashWriter, v *int64) {
|
||||||
|
if v == nil {
|
||||||
|
_, _ = h.Write([]byte{0})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_, _ = h.Write([]byte{1})
|
||||||
|
writeInt64(h, *v)
|
||||||
|
}
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"unicode/utf8"
|
||||||
|
)
|
||||||
|
|
||||||
|
func isEndpointIDChar(c rune) bool {
|
||||||
|
return (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '_' || c == '.' || c == '-'
|
||||||
|
}
|
||||||
|
|
||||||
|
func isMessageIDChar(c rune) bool {
|
||||||
|
return (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '_' || c == '.' || c == '-'
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidEndpointID 校验端编号或群编号(小写字母、数字、_ . -,1–64)。
|
||||||
|
func ValidEndpointID(id string) bool {
|
||||||
|
n := len(id)
|
||||||
|
if n < MinIDLen || n > MaxIDLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, c := range id {
|
||||||
|
if !isEndpointIDChar(c) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidMessageID 校验消息号(字母数字 _ . -,区分大小写,1–64)。
|
||||||
|
func ValidMessageID(id string) bool {
|
||||||
|
n := len(id)
|
||||||
|
if n < MinIDLen || n > MaxIDLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, c := range id {
|
||||||
|
if !isMessageIDChar(c) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidLoginPassword 校验登录密码:8–128 字符,且不能以 nst_ 开头。空串表示由服务器生成,视为合法。
|
||||||
|
func ValidLoginPassword(pw string) bool {
|
||||||
|
if pw == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(pw, SessionTokenPrefix) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
n := utf8.RuneCountInString(pw)
|
||||||
|
return n >= MinLoginPasswordLen && n <= MaxLoginPasswordLen
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoginPasswordForbiddenPrefix 报告密码是否因 nst_ 前缀非法。
|
||||||
|
func LoginPasswordForbiddenPrefix(pw string) bool {
|
||||||
|
return strings.HasPrefix(pw, SessionTokenPrefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidTalkPassword 校验对话密码:空表示清除/不设;否则 4–64 字符。
|
||||||
|
func ValidTalkPassword(pw string) bool {
|
||||||
|
if pw == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
n := utf8.RuneCountInString(pw)
|
||||||
|
return n >= MinTalkPasswordLen && n <= MaxTalkPasswordLen
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidName 校验名称:最多 64 个 Unicode 字符。
|
||||||
|
func ValidName(name string) bool {
|
||||||
|
return utf8.RuneCountInString(name) <= MaxNameChars
|
||||||
|
}
|
||||||
|
|
||||||
|
func requireRID(rid string) error {
|
||||||
|
if rid == "" {
|
||||||
|
return badRequest("missing rid")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func requireVersion(v int) error {
|
||||||
|
if v != Version {
|
||||||
|
return badRequest("unsupported version")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func requireType(got, want string) error {
|
||||||
|
if got != want {
|
||||||
|
return badRequest("wrong type")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,535 @@
|
|||||||
|
package protocol_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
)
|
||||||
|
|
||||||
|
func ptrInt(v int) *int { return &v }
|
||||||
|
func ptrInt64(v int64) *int64 { return &v }
|
||||||
|
func ptrBool(v bool) *bool { return &v }
|
||||||
|
|
||||||
|
func mustMarshal(t *testing.T, v any) []byte {
|
||||||
|
t.Helper()
|
||||||
|
b, err := protocol.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal: %v", err)
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEncodeNoHTMLEscapeAndNoNewline(t *testing.T) {
|
||||||
|
msg := &protocol.Msg{
|
||||||
|
V: protocol.Version,
|
||||||
|
Type: protocol.TypeMsg,
|
||||||
|
ID: "id-1",
|
||||||
|
From: "app-1",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "你好<script>&"},
|
||||||
|
}
|
||||||
|
b := mustMarshal(t, msg)
|
||||||
|
if bytes.Contains(b, []byte(`\u`)) {
|
||||||
|
t.Fatalf("unexpected unicode escape: %s", b)
|
||||||
|
}
|
||||||
|
if bytes.Contains(b, []byte(`\u003c`)) || !bytes.Contains(b, []byte("<script>")) {
|
||||||
|
t.Fatalf("HTML should not be escaped: %s", b)
|
||||||
|
}
|
||||||
|
if bytes.Contains(b, []byte("你好")) == false {
|
||||||
|
t.Fatalf("non-ASCII should remain: %s", b)
|
||||||
|
}
|
||||||
|
if len(b) == 0 || b[len(b)-1] == '\n' {
|
||||||
|
t.Fatalf("must not end with newline: %q", b)
|
||||||
|
}
|
||||||
|
n, err := protocol.FrameBytes(msg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != len(b) {
|
||||||
|
t.Fatalf("FrameBytes=%d len=%d", n, len(b))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeRoundTripAllFrames(t *testing.T) {
|
||||||
|
lim := protocol.DefaultLimits()
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
in any
|
||||||
|
typ string
|
||||||
|
val func(any) error
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "hello",
|
||||||
|
in: &protocol.Hello{
|
||||||
|
V: protocol.Version, Type: protocol.TypeHello, RID: "1",
|
||||||
|
MaxReceiveBytes: ptrInt(4096), Client: "go-sdk/0.1",
|
||||||
|
},
|
||||||
|
typ: protocol.TypeHello,
|
||||||
|
val: func(v any) error { return v.(*protocol.Hello).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "resp_ok",
|
||||||
|
in: &protocol.Resp{
|
||||||
|
V: protocol.Version, Type: protocol.TypeResp, RID: "1", OK: true,
|
||||||
|
Data: protocol.MustRaw(protocol.HelloData{
|
||||||
|
ServerTimeMs: 1, ServerVersion: "0.1.0",
|
||||||
|
MaxBodyBytes: 262144, MaxMetaBytes: 4096, MaxFrameBytes: 786432,
|
||||||
|
MaxTTLSeconds: 1, MaxScheduleSeconds: 1, AckTimeoutSeconds: 1,
|
||||||
|
SessionToken: "nst_abc",
|
||||||
|
}),
|
||||||
|
},
|
||||||
|
typ: protocol.TypeResp,
|
||||||
|
val: func(v any) error { return v.(*protocol.Resp).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "resp_err",
|
||||||
|
in: &protocol.Resp{
|
||||||
|
V: protocol.Version, Type: protocol.TypeResp, RID: "1", OK: false,
|
||||||
|
Error: &protocol.ErrorBody{Code: protocol.CodeTalkPasswordRequired, Message: "需要对话密码"},
|
||||||
|
},
|
||||||
|
typ: protocol.TypeResp,
|
||||||
|
val: func(v any) error { return v.(*protocol.Resp).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "send",
|
||||||
|
in: &protocol.Send{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSend, RID: "2", ID: "018f-Ab",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"},
|
||||||
|
Meta: map[string]any{"k": "v"},
|
||||||
|
},
|
||||||
|
typ: protocol.TypeSend,
|
||||||
|
val: func(v any) error { return v.(*protocol.Send).Validate(lim) },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "msg",
|
||||||
|
in: &protocol.Msg{
|
||||||
|
V: protocol.Version, Type: protocol.TypeMsg, ID: "018f",
|
||||||
|
From: "app-1", To: protocol.Target{Kind: protocol.TargetGroup, ID: "g_ab12cd34"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"}, SendAtMs: 1,
|
||||||
|
},
|
||||||
|
typ: protocol.TypeMsg,
|
||||||
|
val: func(v any) error { return v.(*protocol.Msg).Validate(lim) },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ack",
|
||||||
|
in: &protocol.Ack{V: protocol.Version, Type: protocol.TypeAck, RID: "3", From: "app-1", ID: "018f"},
|
||||||
|
typ: protocol.TypeAck,
|
||||||
|
val: func(v any) error { return v.(*protocol.Ack).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "recall",
|
||||||
|
in: &protocol.Recall{V: protocol.Version, Type: protocol.TypeRecall, RID: "4", ID: "018f"},
|
||||||
|
typ: protocol.TypeRecall,
|
||||||
|
val: func(v any) error { return v.(*protocol.Recall).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "status",
|
||||||
|
in: &protocol.Status{V: protocol.Version, Type: protocol.TypeStatus, RID: "5", ID: "018f", Limit: 100},
|
||||||
|
typ: protocol.TypeStatus,
|
||||||
|
val: func(v any) error { return v.(*protocol.Status).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "receipt",
|
||||||
|
in: &protocol.Receipt{
|
||||||
|
V: protocol.Version, Type: protocol.TypeReceipt, ReceiptID: "9001",
|
||||||
|
ID: "018f", EndpointID: "device-1", State: "accepted", AtMs: 1,
|
||||||
|
},
|
||||||
|
typ: protocol.TypeReceipt,
|
||||||
|
val: func(v any) error { return v.(*protocol.Receipt).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "receipt_ack",
|
||||||
|
in: &protocol.ReceiptAck{V: protocol.Version, Type: protocol.TypeReceiptAck, RID: "6", ReceiptID: "9001"},
|
||||||
|
typ: protocol.TypeReceiptAck,
|
||||||
|
val: func(v any) error { return v.(*protocol.ReceiptAck).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "revoked",
|
||||||
|
in: &protocol.Revoked{V: protocol.Version, Type: protocol.TypeRevoked, ID: "018f", From: "app-1", Reason: "recalled"},
|
||||||
|
typ: protocol.TypeRevoked,
|
||||||
|
val: func(v any) error { return v.(*protocol.Revoked).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "presence.get",
|
||||||
|
in: &protocol.PresenceGet{V: protocol.Version, Type: protocol.TypePresenceGet, RID: "7", IDs: []string{"a", "b"}},
|
||||||
|
typ: protocol.TypePresenceGet,
|
||||||
|
val: func(v any) error { return v.(*protocol.PresenceGet).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "directory.list",
|
||||||
|
in: &protocol.DirectoryList{V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "8", Limit: 100},
|
||||||
|
typ: protocol.TypeDirectoryList,
|
||||||
|
val: func(v any) error { return v.(*protocol.DirectoryList).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "presence.watch",
|
||||||
|
in: &protocol.PresenceWatch{V: protocol.Version, Type: protocol.TypePresenceWatch, RID: "9", IDs: []string{"a"}, All: false},
|
||||||
|
typ: protocol.TypePresenceWatch,
|
||||||
|
val: func(v any) error { return v.(*protocol.PresenceWatch).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "presence",
|
||||||
|
in: &protocol.Presence{V: protocol.Version, Type: protocol.TypePresence, ID: "a", Online: true, AtMs: 1},
|
||||||
|
typ: protocol.TypePresence,
|
||||||
|
val: func(v any) error { return v.(*protocol.Presence).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unlock",
|
||||||
|
in: &protocol.Unlock{V: protocol.Version, Type: protocol.TypeUnlock, RID: "10", EndpointID: "b", TalkPassword: "secret"},
|
||||||
|
typ: protocol.TypeUnlock,
|
||||||
|
val: func(v any) error { return v.(*protocol.Unlock).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "self.get",
|
||||||
|
in: &protocol.SelfGet{V: protocol.Version, Type: protocol.TypeSelfGet, RID: "11"},
|
||||||
|
typ: protocol.TypeSelfGet,
|
||||||
|
val: func(v any) error { return v.(*protocol.SelfGet).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "self.update",
|
||||||
|
in: &protocol.SelfUpdate{V: protocol.Version, Type: protocol.TypeSelfUpdate, RID: "12", Name: "门口", DefaultDelayMs: ptrInt64(10000)},
|
||||||
|
typ: protocol.TypeSelfUpdate,
|
||||||
|
val: func(v any) error { return v.(*protocol.SelfUpdate).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "self.talk_password",
|
||||||
|
in: &protocol.SelfTalkPassword{V: protocol.Version, Type: protocol.TypeSelfTalkPassword, RID: "13", TalkPassword: ""},
|
||||||
|
typ: protocol.TypeSelfTalkPassword,
|
||||||
|
val: func(v any) error { return v.(*protocol.SelfTalkPassword).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "self.login_password",
|
||||||
|
in: &protocol.SelfLoginPassword{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSelfLoginPassword, RID: "14",
|
||||||
|
OldPassword: "oldpass12", NewPassword: "newpass12",
|
||||||
|
},
|
||||||
|
typ: protocol.TypeSelfLoginPassword,
|
||||||
|
val: func(v any) error { return v.(*protocol.SelfLoginPassword).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "self.logout",
|
||||||
|
in: &protocol.SelfLogout{V: protocol.Version, Type: protocol.TypeSelfLogout, RID: "24"},
|
||||||
|
typ: protocol.TypeSelfLogout,
|
||||||
|
val: func(v any) error { return v.(*protocol.SelfLogout).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.create",
|
||||||
|
in: &protocol.GroupCreate{
|
||||||
|
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "15", ID: "",
|
||||||
|
Name: "一组", Members: []protocol.GroupMemberIn{{ID: "b", TalkPassword: "secret"}},
|
||||||
|
},
|
||||||
|
typ: protocol.TypeGroupCreate,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupCreate).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.add",
|
||||||
|
in: &protocol.GroupAdd{
|
||||||
|
V: protocol.Version, Type: protocol.TypeGroupAdd, RID: "16", GroupID: "g_ab12cd34",
|
||||||
|
Members: []protocol.GroupMemberIn{{ID: "c"}},
|
||||||
|
},
|
||||||
|
typ: protocol.TypeGroupAdd,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupAdd).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.remove",
|
||||||
|
in: &protocol.GroupRemove{V: protocol.Version, Type: protocol.TypeGroupRemove, RID: "17", GroupID: "g_ab12cd34", EndpointID: "c"},
|
||||||
|
typ: protocol.TypeGroupRemove,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupRemove).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.leave",
|
||||||
|
in: &protocol.GroupLeave{V: protocol.Version, Type: protocol.TypeGroupLeave, RID: "18", GroupID: "g_ab12cd34"},
|
||||||
|
typ: protocol.TypeGroupLeave,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupLeave).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.transfer",
|
||||||
|
in: &protocol.GroupTransfer{V: protocol.Version, Type: protocol.TypeGroupTransfer, RID: "19", GroupID: "g_ab12cd34", EndpointID: "b"},
|
||||||
|
typ: protocol.TypeGroupTransfer,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupTransfer).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.rename",
|
||||||
|
in: &protocol.GroupRename{V: protocol.Version, Type: protocol.TypeGroupRename, RID: "20", GroupID: "g_ab12cd34", Name: "新名"},
|
||||||
|
typ: protocol.TypeGroupRename,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupRename).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.dissolve",
|
||||||
|
in: &protocol.GroupDissolve{V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "21", GroupID: "g_ab12cd34"},
|
||||||
|
typ: protocol.TypeGroupDissolve,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupDissolve).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.list",
|
||||||
|
in: &protocol.GroupList{V: protocol.Version, Type: protocol.TypeGroupList, RID: "22", Limit: 100},
|
||||||
|
typ: protocol.TypeGroupList,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupList).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group.get",
|
||||||
|
in: &protocol.GroupGet{V: protocol.Version, Type: protocol.TypeGroupGet, RID: "23", GroupID: "g_ab12cd34", Limit: 100},
|
||||||
|
typ: protocol.TypeGroupGet,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupGet).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "group_event",
|
||||||
|
in: &protocol.GroupEvent{
|
||||||
|
V: protocol.Version, Type: protocol.TypeGroupEvent, GroupID: "g_ab12cd34",
|
||||||
|
Event: "member_added", EndpointID: "c", AtMs: 1,
|
||||||
|
},
|
||||||
|
typ: protocol.TypeGroupEvent,
|
||||||
|
val: func(v any) error { return v.(*protocol.GroupEvent).Validate() },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "fatal",
|
||||||
|
in: &protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: "disabled"},
|
||||||
|
typ: protocol.TypeFatal,
|
||||||
|
val: func(v any) error { return v.(*protocol.Fatal).Validate() },
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
raw := mustMarshal(t, tc.in)
|
||||||
|
got, err := protocol.Decode(raw)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Decode: %v", err)
|
||||||
|
}
|
||||||
|
back := mustMarshal(t, got)
|
||||||
|
var a, b any
|
||||||
|
if err := json.Unmarshal(raw, &a); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(back, &b); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
aj, _ := json.Marshal(a)
|
||||||
|
bj, _ := json.Marshal(b)
|
||||||
|
if !bytes.Equal(aj, bj) {
|
||||||
|
t.Fatalf("round-trip mismatch\n%s\n%s", aj, bj)
|
||||||
|
}
|
||||||
|
peek := struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
}{}
|
||||||
|
_ = json.Unmarshal(raw, &peek)
|
||||||
|
if peek.Type != tc.typ {
|
||||||
|
t.Fatalf("type=%s want %s", peek.Type, tc.typ)
|
||||||
|
}
|
||||||
|
if err := tc.val(got); err != nil {
|
||||||
|
t.Fatalf("Validate: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegisterRoundTrip(t *testing.T) {
|
||||||
|
req := &protocol.RegisterRequest{
|
||||||
|
RegistrationCode: "code",
|
||||||
|
ID: "device-1",
|
||||||
|
LoginPassword: "password1",
|
||||||
|
Name: "门口",
|
||||||
|
TalkPassword: "talk",
|
||||||
|
}
|
||||||
|
raw := mustMarshal(t, req)
|
||||||
|
got, err := protocol.DecodeRegister(raw)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := got.Validate(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp := protocol.RegisterResponse{OK: true, Data: protocol.RegisterData{ID: "e_ab12cd34", LoginPassword: "generated"}}
|
||||||
|
rb := mustMarshal(t, resp)
|
||||||
|
var decoded protocol.RegisterResponse
|
||||||
|
if err := protocol.Unmarshal(rb, &decoded); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !decoded.OK || decoded.Data.ID != "e_ab12cd34" {
|
||||||
|
t.Fatalf("unexpected resp: %+v", decoded)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateTable(t *testing.T) {
|
||||||
|
lim := protocol.DefaultLimits()
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
check func() error
|
||||||
|
wantCode string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "endpoint_id_uppercase",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.Send{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "Device"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
|
||||||
|
}).Validate(lim)
|
||||||
|
},
|
||||||
|
wantCode: protocol.CodeBadRequest,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "message_id_allows_upper",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.Send{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "Msg_1.A-b",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
|
||||||
|
}).Validate(lim)
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "send_at_and_delay_mutex",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.Send{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
|
||||||
|
SendAtMs: ptrInt64(1), DelayMs: ptrInt64(2),
|
||||||
|
}).Validate(lim)
|
||||||
|
},
|
||||||
|
wantCode: protocol.CodeBadRequest,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "body_too_large",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.Send{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: strings.Repeat("a", 10)},
|
||||||
|
}).Validate(protocol.Limits{MaxBodyBytes: 5, MaxMetaBytes: 4096, MaxFrameBytes: 786432})
|
||||||
|
},
|
||||||
|
wantCode: protocol.CodeBodyTooLarge,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "meta_too_large",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.Send{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSend, RID: "1", ID: "m1",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "x"},
|
||||||
|
Meta: map[string]any{"k": strings.Repeat("v", 100)},
|
||||||
|
}).Validate(protocol.Limits{MaxBodyBytes: 262144, MaxMetaBytes: 20, MaxFrameBytes: 786432})
|
||||||
|
},
|
||||||
|
wantCode: protocol.CodeMetaTooLarge,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "login_password_nst_prefix",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.RegisterRequest{LoginPassword: "nst_notallowed"}).Validate()
|
||||||
|
},
|
||||||
|
wantCode: protocol.CodeBadRequest,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "self_login_password_nst_prefix",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.SelfLoginPassword{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSelfLoginPassword, RID: "1",
|
||||||
|
OldPassword: "oldpass12", NewPassword: "nst_tokenlike",
|
||||||
|
}).Validate()
|
||||||
|
},
|
||||||
|
wantCode: protocol.CodeBadRequest,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "hello_max_receive_too_small",
|
||||||
|
check: func() error {
|
||||||
|
return (&protocol.Hello{
|
||||||
|
V: protocol.Version, Type: protocol.TypeHello, RID: "1",
|
||||||
|
MaxReceiveBytes: ptrInt(100),
|
||||||
|
}).Validate()
|
||||||
|
},
|
||||||
|
wantCode: protocol.CodeBadRequest,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
err := tc.check()
|
||||||
|
if tc.wantCode == "" {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected err: %v", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
pe, ok := err.(*protocol.Error)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("want *protocol.Error, got %T %v", err, err)
|
||||||
|
}
|
||||||
|
if pe.Code != tc.wantCode {
|
||||||
|
t.Fatalf("code=%s want %s", pe.Code, tc.wantCode)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequestFingerprintMetaKeyOrderInsensitive(t *testing.T) {
|
||||||
|
base := func(meta map[string]any) *protocol.Send {
|
||||||
|
return &protocol.Send{
|
||||||
|
V: protocol.Version, Type: protocol.TypeSend, RID: "rid-ignored", ID: "m1",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: "device-1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, ContentType: "text/plain", Data: "hello"},
|
||||||
|
Meta: meta,
|
||||||
|
Offline: &protocol.OfflineOpts{Keep: true, TTLSeconds: ptrInt64(60)},
|
||||||
|
Receipt: ptrBool(true),
|
||||||
|
TalkPassword: "should-not-affect",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fp1, err := protocol.RequestFingerprint(base(map[string]any{"a": 1, "b": "x", "c": true}))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
fp2, err := protocol.RequestFingerprint(base(map[string]any{"c": true, "b": "x", "a": 1}))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if fp1 != fp2 {
|
||||||
|
t.Fatalf("fingerprint depends on meta key order: %s vs %s", fp1, fp2)
|
||||||
|
}
|
||||||
|
// rid / talk_password 变化不应影响
|
||||||
|
s3 := base(map[string]any{"a": 1, "b": "x", "c": true})
|
||||||
|
s3.RID = "other-rid"
|
||||||
|
s3.TalkPassword = "other"
|
||||||
|
fp3, err := protocol.RequestFingerprint(s3)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if fp3 != fp1 {
|
||||||
|
t.Fatalf("rid/talk_password affected fingerprint")
|
||||||
|
}
|
||||||
|
// 正文变化应影响
|
||||||
|
s4 := base(map[string]any{"a": 1, "b": "x", "c": true})
|
||||||
|
s4.Body.Data = "hello!"
|
||||||
|
fp4, err := protocol.RequestFingerprint(s4)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if fp4 == fp1 {
|
||||||
|
t.Fatalf("body change should change fingerprint")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestErrorCodesDefined(t *testing.T) {
|
||||||
|
codes := []string{
|
||||||
|
protocol.CodeBadRequest, protocol.CodeNotReady, protocol.CodeUnauthorized,
|
||||||
|
protocol.CodeForbidden, protocol.CodeNotFound, protocol.CodeInvalidTarget,
|
||||||
|
protocol.CodeConflict, protocol.CodeIDTaken, protocol.CodeBodyTooLarge,
|
||||||
|
protocol.CodeMetaTooLarge, protocol.CodeFrameTooLarge, protocol.CodeResponseTooLarge,
|
||||||
|
protocol.CodeTalkPasswordRequired, protocol.CodeTalkPasswordInvalid,
|
||||||
|
protocol.CodeRateLimited, protocol.CodeNotMember, protocol.CodeOwnerCannotLeave,
|
||||||
|
protocol.CodeGroupFull, protocol.CodeQuotaExceeded, protocol.CodeEndpointDisabled,
|
||||||
|
protocol.CodeRegistrationClosed, protocol.CodeRegistrationCodeInvalid, protocol.CodeBusy,
|
||||||
|
}
|
||||||
|
if len(codes) != 23 {
|
||||||
|
t.Fatalf("want 23 error codes, got %d", len(codes))
|
||||||
|
}
|
||||||
|
seen := map[string]bool{}
|
||||||
|
for _, c := range codes {
|
||||||
|
if c == "" || seen[c] {
|
||||||
|
t.Fatalf("bad code %q", c)
|
||||||
|
}
|
||||||
|
seen[c] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,393 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
import "encoding/json"
|
||||||
|
|
||||||
|
// Target 是发送目标。
|
||||||
|
type Target struct {
|
||||||
|
Kind string `json:"kind"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Body 是消息正文。
|
||||||
|
type Body struct {
|
||||||
|
Enc string `json:"enc"`
|
||||||
|
ContentType string `json:"content_type,omitempty"`
|
||||||
|
Data string `json:"data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// OfflineOpts 是离线保留选项。
|
||||||
|
type OfflineOpts struct {
|
||||||
|
Keep bool `json:"keep"`
|
||||||
|
TTLSeconds *int64 `json:"ttl_seconds,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Limits 是与服务器配置相关的校验上限。
|
||||||
|
type Limits struct {
|
||||||
|
MaxBodyBytes int
|
||||||
|
MaxMetaBytes int
|
||||||
|
MaxFrameBytes int
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultLimits 返回 DEVELOPMENT 示例中的默认上限。
|
||||||
|
func DefaultLimits() Limits {
|
||||||
|
return Limits{
|
||||||
|
MaxBodyBytes: DefaultMaxBodyBytes,
|
||||||
|
MaxMetaBytes: DefaultMaxMetaBytes,
|
||||||
|
MaxFrameBytes: DefaultMaxFrameBytes,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l Limits) withDefaults() Limits {
|
||||||
|
if l.MaxBodyBytes <= 0 {
|
||||||
|
l.MaxBodyBytes = DefaultMaxBodyBytes
|
||||||
|
}
|
||||||
|
if l.MaxMetaBytes <= 0 {
|
||||||
|
l.MaxMetaBytes = DefaultMaxMetaBytes
|
||||||
|
}
|
||||||
|
if l.MaxFrameBytes <= 0 {
|
||||||
|
l.MaxFrameBytes = DefaultMaxFrameBytes
|
||||||
|
}
|
||||||
|
return l
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hello 是握手请求(第 6.1 节)。
|
||||||
|
type Hello struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
MaxReceiveBytes *int `json:"max_receive_bytes,omitempty"`
|
||||||
|
Client string `json:"client,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// HelloData 是握手成功响应 data。
|
||||||
|
type HelloData struct {
|
||||||
|
ServerTimeMs int64 `json:"server_time_ms"`
|
||||||
|
ServerVersion string `json:"server_version"`
|
||||||
|
MaxBodyBytes int `json:"max_body_bytes"`
|
||||||
|
MaxMetaBytes int `json:"max_meta_bytes"`
|
||||||
|
MaxFrameBytes int `json:"max_frame_bytes"`
|
||||||
|
MaxTTLSeconds int64 `json:"max_ttl_seconds"`
|
||||||
|
MaxScheduleSeconds int64 `json:"max_schedule_seconds"`
|
||||||
|
AckTimeoutSeconds int64 `json:"ack_timeout_seconds"`
|
||||||
|
SessionToken string `json:"session_token,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Resp 是通用响应帧。
|
||||||
|
type Resp struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
OK bool `json:"ok"`
|
||||||
|
Data json.RawMessage `json:"data,omitempty"`
|
||||||
|
Error *ErrorBody `json:"error,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send 是发送请求(第 6.2 节)。
|
||||||
|
type Send struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
To Target `json:"to"`
|
||||||
|
Body Body `json:"body"`
|
||||||
|
Meta map[string]any `json:"meta,omitempty"`
|
||||||
|
DelayMs *int64 `json:"delay_ms,omitempty"`
|
||||||
|
SendAtMs *int64 `json:"send_at_ms,omitempty"`
|
||||||
|
Offline *OfflineOpts `json:"offline,omitempty"`
|
||||||
|
Receipt *bool `json:"receipt,omitempty"`
|
||||||
|
TalkPassword string `json:"talk_password,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendData 是发送成功响应 data。
|
||||||
|
type SendData struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
SendAtMs int64 `json:"send_at_ms"`
|
||||||
|
State string `json:"state"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Msg 是下行消息(第 6.3 节)。
|
||||||
|
type Msg struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
From string `json:"from"`
|
||||||
|
To Target `json:"to"`
|
||||||
|
Body Body `json:"body"`
|
||||||
|
Meta map[string]any `json:"meta,omitempty"`
|
||||||
|
SendAtMs int64 `json:"send_at_ms"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ack 是确认请求。
|
||||||
|
type Ack struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
From string `json:"from"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Recall 是撤回请求。
|
||||||
|
type Recall struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecallData 是撤回响应 data。
|
||||||
|
type RecallData struct {
|
||||||
|
Result string `json:"result"`
|
||||||
|
Recalled int `json:"recalled"`
|
||||||
|
Accepted int `json:"accepted"`
|
||||||
|
Other int `json:"other"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Status 是状态查询请求。
|
||||||
|
type Status struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
Cursor string `json:"cursor,omitempty"`
|
||||||
|
Limit int `json:"limit,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Receipt 是回执下行。
|
||||||
|
type Receipt struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
ReceiptID string `json:"receipt_id"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
EndpointID string `json:"endpoint_id"`
|
||||||
|
State string `json:"state"`
|
||||||
|
Reason string `json:"reason"`
|
||||||
|
AtMs int64 `json:"at_ms"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReceiptAck 是回执确认。
|
||||||
|
type ReceiptAck struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
ReceiptID string `json:"receipt_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Revoked 是已推送消息作废通知。
|
||||||
|
type Revoked struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
From string `json:"from"`
|
||||||
|
Reason string `json:"reason"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PresenceGet 查询在线状态。
|
||||||
|
type PresenceGet struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
IDs []string `json:"ids"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// DirectoryList 列目录。
|
||||||
|
type DirectoryList struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
Cursor string `json:"cursor,omitempty"`
|
||||||
|
Limit int `json:"limit,omitempty"`
|
||||||
|
Query string `json:"query,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PresenceWatch 订阅上下线。
|
||||||
|
type PresenceWatch struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
IDs []string `json:"ids,omitempty"`
|
||||||
|
All bool `json:"all,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Presence 是上下线通知。
|
||||||
|
type Presence struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
Online bool `json:"online"`
|
||||||
|
AtMs int64 `json:"at_ms"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unlock 解锁对话密码。
|
||||||
|
type Unlock struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
EndpointID string `json:"endpoint_id"`
|
||||||
|
TalkPassword string `json:"talk_password"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelfGet 读取自己的资料。
|
||||||
|
type SelfGet struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelfUpdate 更新自己的资料。
|
||||||
|
type SelfUpdate struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
Name string `json:"name,omitempty"`
|
||||||
|
DefaultDelayMs *int64 `json:"default_delay_ms,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelfTalkPassword 设置对话密码。
|
||||||
|
type SelfTalkPassword struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
TalkPassword string `json:"talk_password"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelfLoginPassword 修改登录密码。
|
||||||
|
type SelfLoginPassword struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
OldPassword string `json:"old_password"`
|
||||||
|
NewPassword string `json:"new_password"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelfLogout 退出登录并作废会话令牌。
|
||||||
|
type SelfLogout struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupMemberIn 是建群/加群时的成员项。
|
||||||
|
type GroupMemberIn struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
TalkPassword string `json:"talk_password,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupCreate 建群。
|
||||||
|
type GroupCreate struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Members []GroupMemberIn `json:"members"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupAdd 加成员。
|
||||||
|
type GroupAdd struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
Members []GroupMemberIn `json:"members"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupRemove 移除成员。
|
||||||
|
type GroupRemove struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
EndpointID string `json:"endpoint_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupLeave 退群。
|
||||||
|
type GroupLeave struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupTransfer 转让群主。
|
||||||
|
type GroupTransfer struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
EndpointID string `json:"endpoint_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupRename 改群名。
|
||||||
|
type GroupRename struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupDissolve 解散群。
|
||||||
|
type GroupDissolve struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupList 列出我的群。
|
||||||
|
type GroupList struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
Cursor string `json:"cursor,omitempty"`
|
||||||
|
Limit int `json:"limit,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupGet 获取群详情。
|
||||||
|
type GroupGet struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
RID string `json:"rid"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
Cursor string `json:"cursor,omitempty"`
|
||||||
|
Limit int `json:"limit,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GroupEvent 是群事件下行。
|
||||||
|
type GroupEvent struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
GroupID string `json:"group_id"`
|
||||||
|
Event string `json:"event"`
|
||||||
|
EndpointID string `json:"endpoint_id,omitempty"`
|
||||||
|
AtMs int64 `json:"at_ms"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fatal 是致命错误下行。
|
||||||
|
type Fatal struct {
|
||||||
|
V int `json:"v"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
Reason string `json:"reason"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterRequest 是 HTTP 注册请求(第 6.9 节)。
|
||||||
|
type RegisterRequest struct {
|
||||||
|
RegistrationCode string `json:"registration_code"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
LoginPassword string `json:"login_password"`
|
||||||
|
Name string `json:"name,omitempty"`
|
||||||
|
TalkPassword string `json:"talk_password,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterData 是注册成功 data。
|
||||||
|
type RegisterData struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
LoginPassword string `json:"login_password,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterResponse 是注册 HTTP 响应。
|
||||||
|
type RegisterResponse struct {
|
||||||
|
OK bool `json:"ok"`
|
||||||
|
Data RegisterData `json:"data,omitempty"`
|
||||||
|
Error *ErrorBody `json:"error,omitempty"`
|
||||||
|
}
|
||||||
@@ -0,0 +1,750 @@
|
|||||||
|
package protocol
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"unicode/utf8"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DecodeBody 按 enc 解码正文,返回原始字节。
|
||||||
|
func DecodeBody(b Body) ([]byte, error) {
|
||||||
|
switch b.Enc {
|
||||||
|
case EncUTF8:
|
||||||
|
if !utf8.ValidString(b.Data) {
|
||||||
|
return nil, badRequest("body data is not valid utf8")
|
||||||
|
}
|
||||||
|
return []byte(b.Data), nil
|
||||||
|
case EncBase64:
|
||||||
|
raw, err := base64.StdEncoding.DecodeString(b.Data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, badRequest("body data is not valid base64")
|
||||||
|
}
|
||||||
|
return raw, nil
|
||||||
|
default:
|
||||||
|
return nil, badRequest("invalid body.enc")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// EffectiveContentType 返回带默认值的 content_type。
|
||||||
|
func EffectiveContentType(b Body) string {
|
||||||
|
if b.ContentType != "" {
|
||||||
|
return b.ContentType
|
||||||
|
}
|
||||||
|
switch b.Enc {
|
||||||
|
case EncBase64:
|
||||||
|
return DefaultContentTypeBase64
|
||||||
|
default:
|
||||||
|
return DefaultContentTypeUTF8
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// EffectiveReceipt 返回 receipt 有效值(默认 true)。
|
||||||
|
func EffectiveReceipt(s *Send) bool {
|
||||||
|
if s.Receipt == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return *s.Receipt
|
||||||
|
}
|
||||||
|
|
||||||
|
// EffectiveOfflineKeep 返回 offline.keep(默认 false)。
|
||||||
|
func EffectiveOfflineKeep(s *Send) bool {
|
||||||
|
if s.Offline == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return s.Offline.Keep
|
||||||
|
}
|
||||||
|
|
||||||
|
// EffectiveOfflineTTL 返回 offline.ttl_seconds 有效值;keep 为 false 时为 0。
|
||||||
|
func EffectiveOfflineTTL(s *Send) int64 {
|
||||||
|
if !EffectiveOfflineKeep(s) {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if s.Offline.TTLSeconds == nil {
|
||||||
|
return DefaultOfflineTTL
|
||||||
|
}
|
||||||
|
return *s.Offline.TTLSeconds
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkMetaSize(meta map[string]any, maxBytes int) error {
|
||||||
|
if meta == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
b, err := Marshal(meta)
|
||||||
|
if err != nil {
|
||||||
|
return badRequest("invalid meta")
|
||||||
|
}
|
||||||
|
if len(b) > maxBytes {
|
||||||
|
return metaTooLarge("meta exceeds limit")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkBodySize(b Body, maxBytes int) ([]byte, error) {
|
||||||
|
raw, err := DecodeBody(b)
|
||||||
|
if err != nil {
|
||||||
|
if pe, ok := err.(*Error); ok {
|
||||||
|
return nil, pe
|
||||||
|
}
|
||||||
|
return nil, badRequest(err.Error())
|
||||||
|
}
|
||||||
|
if len(raw) > maxBytes {
|
||||||
|
return nil, bodyTooLarge("body exceeds limit")
|
||||||
|
}
|
||||||
|
if b.ContentType != "" && utf8.RuneCountInString(b.ContentType) > MaxContentTypeLen {
|
||||||
|
return nil, badRequest("content_type too long")
|
||||||
|
}
|
||||||
|
return raw, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkPageLimit(limit int) error {
|
||||||
|
if limit < 0 || limit > MaxPageLimit {
|
||||||
|
return badRequest("limit out of range")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Hello。
|
||||||
|
func (h *Hello) Validate() error {
|
||||||
|
if err := requireVersion(h.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(h.Type, TypeHello); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(h.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if h.MaxReceiveBytes != nil && *h.MaxReceiveBytes < MinMaxReceiveBytes {
|
||||||
|
return badRequest("max_receive_bytes too small")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Resp。
|
||||||
|
func (r *Resp) Validate() error {
|
||||||
|
if err := requireVersion(r.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(r.Type, TypeResp); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(r.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !r.OK && (r.Error == nil || r.Error.Code == "") {
|
||||||
|
return badRequest("error response missing code")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Send(含正文/meta 大小与 send_at_ms/delay_ms 互斥)。
|
||||||
|
func (s *Send) Validate(lim Limits) error {
|
||||||
|
lim = lim.withDefaults()
|
||||||
|
if err := requireVersion(s.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(s.Type, TypeSend); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(s.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidMessageID(s.ID) {
|
||||||
|
return badRequest("invalid message id")
|
||||||
|
}
|
||||||
|
if s.To.Kind != TargetEndpoint && s.To.Kind != TargetGroup {
|
||||||
|
return badRequest("invalid to.kind")
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(s.To.ID) {
|
||||||
|
return badRequest("invalid to.id")
|
||||||
|
}
|
||||||
|
if _, err := checkBodySize(s.Body, lim.MaxBodyBytes); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := checkMetaSize(s.Meta, lim.MaxMetaBytes); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if s.SendAtMs != nil && s.DelayMs != nil {
|
||||||
|
return badRequest("send_at_ms and delay_ms are mutually exclusive")
|
||||||
|
}
|
||||||
|
n, err := FrameBytes(s)
|
||||||
|
if err != nil {
|
||||||
|
return badRequest("cannot encode frame")
|
||||||
|
}
|
||||||
|
if n > lim.MaxFrameBytes {
|
||||||
|
return frameTooLarge("frame exceeds limit")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Msg。
|
||||||
|
func (m *Msg) Validate(lim Limits) error {
|
||||||
|
lim = lim.withDefaults()
|
||||||
|
if err := requireVersion(m.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(m.Type, TypeMsg); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidMessageID(m.ID) {
|
||||||
|
return badRequest("invalid message id")
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(m.From) {
|
||||||
|
return badRequest("invalid from")
|
||||||
|
}
|
||||||
|
if m.To.Kind != TargetEndpoint && m.To.Kind != TargetGroup {
|
||||||
|
return badRequest("invalid to.kind")
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(m.To.ID) {
|
||||||
|
return badRequest("invalid to.id")
|
||||||
|
}
|
||||||
|
if _, err := checkBodySize(m.Body, lim.MaxBodyBytes); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := checkMetaSize(m.Meta, lim.MaxMetaBytes); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Ack。
|
||||||
|
func (a *Ack) Validate() error {
|
||||||
|
if err := requireVersion(a.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(a.Type, TypeAck); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(a.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(a.From) {
|
||||||
|
return badRequest("invalid from")
|
||||||
|
}
|
||||||
|
if !ValidMessageID(a.ID) {
|
||||||
|
return badRequest("invalid message id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Recall。
|
||||||
|
func (r *Recall) Validate() error {
|
||||||
|
if err := requireVersion(r.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(r.Type, TypeRecall); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(r.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidMessageID(r.ID) {
|
||||||
|
return badRequest("invalid message id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Status。
|
||||||
|
func (s *Status) Validate() error {
|
||||||
|
if err := requireVersion(s.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(s.Type, TypeStatus); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(s.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidMessageID(s.ID) {
|
||||||
|
return badRequest("invalid message id")
|
||||||
|
}
|
||||||
|
if err := checkPageLimit(s.Limit); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Receipt。
|
||||||
|
func (r *Receipt) Validate() error {
|
||||||
|
if err := requireVersion(r.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(r.Type, TypeReceipt); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if r.ReceiptID == "" {
|
||||||
|
return badRequest("missing receipt_id")
|
||||||
|
}
|
||||||
|
if !ValidMessageID(r.ID) {
|
||||||
|
return badRequest("invalid message id")
|
||||||
|
}
|
||||||
|
if r.EndpointID != "" && !ValidEndpointID(r.EndpointID) {
|
||||||
|
return badRequest("invalid endpoint_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 ReceiptAck。
|
||||||
|
func (r *ReceiptAck) Validate() error {
|
||||||
|
if err := requireVersion(r.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(r.Type, TypeReceiptAck); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(r.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if r.ReceiptID == "" {
|
||||||
|
return badRequest("missing receipt_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Revoked。
|
||||||
|
func (r *Revoked) Validate() error {
|
||||||
|
if err := requireVersion(r.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(r.Type, TypeRevoked); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidMessageID(r.ID) {
|
||||||
|
return badRequest("invalid message id")
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(r.From) {
|
||||||
|
return badRequest("invalid from")
|
||||||
|
}
|
||||||
|
if r.Reason == "" {
|
||||||
|
return badRequest("missing reason")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 PresenceGet。
|
||||||
|
func (p *PresenceGet) Validate() error {
|
||||||
|
if err := requireVersion(p.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(p.Type, TypePresenceGet); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(p.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if len(p.IDs) > MaxPresenceGetIDs {
|
||||||
|
return badRequest("ids count out of range")
|
||||||
|
}
|
||||||
|
for _, id := range p.IDs {
|
||||||
|
if !ValidEndpointID(id) {
|
||||||
|
return badRequest("invalid id in ids")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 DirectoryList。
|
||||||
|
func (d *DirectoryList) Validate() error {
|
||||||
|
if err := requireVersion(d.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(d.Type, TypeDirectoryList); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(d.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := checkPageLimit(d.Limit); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 PresenceWatch。
|
||||||
|
func (p *PresenceWatch) Validate() error {
|
||||||
|
if err := requireVersion(p.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(p.Type, TypePresenceWatch); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(p.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if len(p.IDs) > MaxPresenceWatchIDs {
|
||||||
|
return badRequest("ids count out of range")
|
||||||
|
}
|
||||||
|
for _, id := range p.IDs {
|
||||||
|
if !ValidEndpointID(id) {
|
||||||
|
return badRequest("invalid id in ids")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Presence。
|
||||||
|
func (p *Presence) Validate() error {
|
||||||
|
if err := requireVersion(p.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(p.Type, TypePresence); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(p.ID) {
|
||||||
|
return badRequest("invalid id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Unlock。
|
||||||
|
func (u *Unlock) Validate() error {
|
||||||
|
if err := requireVersion(u.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(u.Type, TypeUnlock); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(u.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(u.EndpointID) {
|
||||||
|
return badRequest("invalid endpoint_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 SelfGet。
|
||||||
|
func (s *SelfGet) Validate() error {
|
||||||
|
if err := requireVersion(s.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(s.Type, TypeSelfGet); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return requireRID(s.RID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 SelfUpdate。
|
||||||
|
func (s *SelfUpdate) Validate() error {
|
||||||
|
if err := requireVersion(s.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(s.Type, TypeSelfUpdate); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(s.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidName(s.Name) {
|
||||||
|
return badRequest("name too long")
|
||||||
|
}
|
||||||
|
if s.DefaultDelayMs != nil && *s.DefaultDelayMs < 0 {
|
||||||
|
return badRequest("default_delay_ms negative")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 SelfTalkPassword。
|
||||||
|
func (s *SelfTalkPassword) Validate() error {
|
||||||
|
if err := requireVersion(s.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(s.Type, TypeSelfTalkPassword); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(s.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidTalkPassword(s.TalkPassword) {
|
||||||
|
return badRequest("invalid talk_password")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 SelfLoginPassword。
|
||||||
|
func (s *SelfLoginPassword) Validate() error {
|
||||||
|
if err := requireVersion(s.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(s.Type, TypeSelfLoginPassword); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(s.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if s.NewPassword == "" {
|
||||||
|
return badRequest("new_password required")
|
||||||
|
}
|
||||||
|
if LoginPasswordForbiddenPrefix(s.NewPassword) {
|
||||||
|
return badRequest("login password must not start with nst_")
|
||||||
|
}
|
||||||
|
if !ValidLoginPassword(s.NewPassword) {
|
||||||
|
return badRequest("invalid new_password")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 SelfLogout。
|
||||||
|
func (s *SelfLogout) Validate() error {
|
||||||
|
if err := requireVersion(s.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(s.Type, TypeSelfLogout); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return requireRID(s.RID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateGroupMembers(members []GroupMemberIn) error {
|
||||||
|
for _, m := range members {
|
||||||
|
if !ValidEndpointID(m.ID) {
|
||||||
|
return badRequest("invalid member id")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupCreate。
|
||||||
|
func (g *GroupCreate) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupCreate); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if g.ID != "" && !ValidEndpointID(g.ID) {
|
||||||
|
return badRequest("invalid group id")
|
||||||
|
}
|
||||||
|
if !ValidName(g.Name) || g.Name == "" {
|
||||||
|
return badRequest("invalid name")
|
||||||
|
}
|
||||||
|
return validateGroupMembers(g.Members)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupAdd。
|
||||||
|
func (g *GroupAdd) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupAdd); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
if len(g.Members) == 0 {
|
||||||
|
return badRequest("members required")
|
||||||
|
}
|
||||||
|
return validateGroupMembers(g.Members)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupRemove。
|
||||||
|
func (g *GroupRemove) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupRemove); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.EndpointID) {
|
||||||
|
return badRequest("invalid endpoint_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupLeave。
|
||||||
|
func (g *GroupLeave) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupLeave); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupTransfer。
|
||||||
|
func (g *GroupTransfer) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupTransfer); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.EndpointID) {
|
||||||
|
return badRequest("invalid endpoint_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupRename。
|
||||||
|
func (g *GroupRename) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupRename); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
if !ValidName(g.Name) || g.Name == "" {
|
||||||
|
return badRequest("invalid name")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupDissolve。
|
||||||
|
func (g *GroupDissolve) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupDissolve); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupList。
|
||||||
|
func (g *GroupList) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupList); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return checkPageLimit(g.Limit)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupGet。
|
||||||
|
func (g *GroupGet) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupGet); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireRID(g.RID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
return checkPageLimit(g.Limit)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 GroupEvent。
|
||||||
|
func (g *GroupEvent) Validate() error {
|
||||||
|
if err := requireVersion(g.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(g.Type, TypeGroupEvent); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ValidEndpointID(g.GroupID) {
|
||||||
|
return badRequest("invalid group_id")
|
||||||
|
}
|
||||||
|
if g.Event == "" {
|
||||||
|
return badRequest("missing event")
|
||||||
|
}
|
||||||
|
if g.EndpointID != "" && !ValidEndpointID(g.EndpointID) {
|
||||||
|
return badRequest("invalid endpoint_id")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 Fatal。
|
||||||
|
func (f *Fatal) Validate() error {
|
||||||
|
if err := requireVersion(f.V); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := requireType(f.Type, TypeFatal); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
switch f.Reason {
|
||||||
|
case "disabled", "deleted", "password_reset", "protocol":
|
||||||
|
return nil
|
||||||
|
default:
|
||||||
|
return badRequest("invalid reason")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 RegisterRequest。
|
||||||
|
func (r *RegisterRequest) Validate() error {
|
||||||
|
if r.ID != "" && !ValidEndpointID(r.ID) {
|
||||||
|
return badRequest("invalid id")
|
||||||
|
}
|
||||||
|
if LoginPasswordForbiddenPrefix(r.LoginPassword) {
|
||||||
|
return badRequest("login password must not start with nst_")
|
||||||
|
}
|
||||||
|
if !ValidLoginPassword(r.LoginPassword) {
|
||||||
|
return badRequest("invalid login_password")
|
||||||
|
}
|
||||||
|
if !ValidName(r.Name) {
|
||||||
|
return badRequest("name too long")
|
||||||
|
}
|
||||||
|
if !ValidTalkPassword(r.TalkPassword) {
|
||||||
|
return badRequest("invalid talk_password")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MetaCanonicalJSON 将 meta 规范化后序列化(键排序;嵌套 map 同样处理)。
|
||||||
|
// encoding/json 对 map[string]T 会按键名字典序输出,因此插入顺序不影响结果。
|
||||||
|
func MetaCanonicalJSON(meta map[string]any) ([]byte, error) {
|
||||||
|
if meta == nil {
|
||||||
|
return []byte("{}"), nil
|
||||||
|
}
|
||||||
|
return Marshal(normalizeMetaValue(meta))
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetaValue(v any) any {
|
||||||
|
switch t := v.(type) {
|
||||||
|
case map[string]any:
|
||||||
|
out := make(map[string]any, len(t))
|
||||||
|
for k, child := range t {
|
||||||
|
out[k] = normalizeMetaValue(child)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case json.Number:
|
||||||
|
return t
|
||||||
|
default:
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
}
|
||||||
+103
-9
@@ -10,29 +10,105 @@ import (
|
|||||||
_ "modernc.org/sqlite"
|
_ "modernc.org/sqlite"
|
||||||
)
|
)
|
||||||
|
|
||||||
// OpenWriter 按 DEVELOPMENT 7.7 打开写连接(带 _txlock=immediate)。
|
const dbFileName = "nixmsg.db"
|
||||||
|
|
||||||
|
// DB 持有读写连接与写入队列。
|
||||||
|
type DB struct {
|
||||||
|
Write *sql.DB
|
||||||
|
Read *sql.DB
|
||||||
|
Queue *Queue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Open 打开数据目录下的库:写连接、读连接池、迁移,并启动简单写入队列。
|
||||||
|
// 若 nixmsg.db 尚不存在则跳过迁移前备份。
|
||||||
|
func Open(dataDir, synchronous string) (*DB, error) {
|
||||||
|
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||||
|
return nil, fmt.Errorf("mkdir data_dir: %w", err)
|
||||||
|
}
|
||||||
|
dbPath := filepath.Join(dataDir, dbFileName)
|
||||||
|
existed, err := fileExists(dbPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
write, err := OpenWriter(dataDir, synchronous)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if migErr := Migrate(write, dataDir, existed); migErr != nil {
|
||||||
|
_ = write.Close()
|
||||||
|
return nil, migErr
|
||||||
|
}
|
||||||
|
read, err := OpenReader(dataDir, synchronous)
|
||||||
|
if err != nil {
|
||||||
|
_ = write.Close()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
q := NewQueue(write)
|
||||||
|
return &DB{Write: write, Read: read, Queue: q}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close 停止写入队列并关闭连接。
|
||||||
|
func (d *DB) Close() error {
|
||||||
|
var first error
|
||||||
|
if d.Queue != nil {
|
||||||
|
if err := d.Queue.Close(); err != nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if d.Read != nil {
|
||||||
|
if err := d.Read.Close(); err != nil && first == nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if d.Write != nil {
|
||||||
|
if err := d.Write.Close(); err != nil && first == nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
// OpenWriter 按 DEVELOPMENT 7.7 打开写连接(带 _txlock=immediate,MaxOpenConns=1)。
|
||||||
func OpenWriter(dataDir, synchronous string) (*sql.DB, error) {
|
func OpenWriter(dataDir, synchronous string) (*sql.DB, error) {
|
||||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||||
return nil, fmt.Errorf("mkdir data_dir: %w", err)
|
return nil, fmt.Errorf("mkdir data_dir: %w", err)
|
||||||
}
|
}
|
||||||
dsn, err := writeDSN(dataDir, synchronous)
|
dsn, err := buildDSN(dataDir, synchronous, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
db, err := sql.Open("sqlite", dsn)
|
db, err := sql.Open("sqlite", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("open sqlite: %w", err)
|
return nil, fmt.Errorf("open sqlite writer: %w", err)
|
||||||
}
|
}
|
||||||
db.SetMaxOpenConns(1)
|
db.SetMaxOpenConns(1)
|
||||||
db.SetMaxIdleConns(1)
|
db.SetMaxIdleConns(1)
|
||||||
if err := db.Ping(); err != nil {
|
if err := db.Ping(); err != nil {
|
||||||
_ = db.Close()
|
_ = db.Close()
|
||||||
return nil, fmt.Errorf("ping sqlite: %w", err)
|
return nil, fmt.Errorf("ping sqlite writer: %w", err)
|
||||||
}
|
}
|
||||||
return db, nil
|
return db, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func writeDSN(dataDir, synchronous string) (string, error) {
|
// OpenReader 打开读连接池(DSN 不含 _txlock=immediate)。
|
||||||
|
func OpenReader(dataDir, synchronous string) (*sql.DB, error) {
|
||||||
|
dsn, err := buildDSN(dataDir, synchronous, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
db, err := sql.Open("sqlite", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("open sqlite reader: %w", err)
|
||||||
|
}
|
||||||
|
if err := db.Ping(); err != nil {
|
||||||
|
_ = db.Close()
|
||||||
|
return nil, fmt.Errorf("ping sqlite reader: %w", err)
|
||||||
|
}
|
||||||
|
return db, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildDSN(dataDir, synchronous string, writer bool) (string, error) {
|
||||||
sync := strings.ToUpper(strings.TrimSpace(synchronous))
|
sync := strings.ToUpper(strings.TrimSpace(synchronous))
|
||||||
if sync == "" {
|
if sync == "" {
|
||||||
sync = "FULL"
|
sync = "FULL"
|
||||||
@@ -40,10 +116,28 @@ func writeDSN(dataDir, synchronous string) (string, error) {
|
|||||||
if sync != "FULL" && sync != "NORMAL" {
|
if sync != "FULL" && sync != "NORMAL" {
|
||||||
return "", fmt.Errorf("invalid sqlite_synchronous: %s", synchronous)
|
return "", fmt.Errorf("invalid sqlite_synchronous: %s", synchronous)
|
||||||
}
|
}
|
||||||
dbPath := filepath.ToSlash(filepath.Join(dataDir, "nixmsg.db"))
|
dbPath := filepath.ToSlash(filepath.Join(dataDir, dbFileName))
|
||||||
return fmt.Sprintf(
|
dsn := fmt.Sprintf(
|
||||||
"file:%s?_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)&_pragma=synchronous(%s)&_pragma=foreign_keys(ON)&_pragma=secure_delete(ON)&_txlock=immediate",
|
"file:%s?_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)&_pragma=synchronous(%s)&_pragma=foreign_keys(ON)&_pragma=secure_delete(ON)",
|
||||||
dbPath,
|
dbPath,
|
||||||
sync,
|
sync,
|
||||||
), nil
|
)
|
||||||
|
if writer {
|
||||||
|
dsn += "&_txlock=immediate"
|
||||||
|
}
|
||||||
|
return dsn, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func fileExists(path string) (bool, error) {
|
||||||
|
st, err := os.Stat(path)
|
||||||
|
if err == nil {
|
||||||
|
if st.IsDir() {
|
||||||
|
return false, fmt.Errorf("db path is a directory: %s", path)
|
||||||
|
}
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
return false, err
|
||||||
}
|
}
|
||||||
|
|||||||
+144
-15
@@ -1,36 +1,165 @@
|
|||||||
package store
|
package store
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestOpenAndMigrate(t *testing.T) {
|
func TestOpenEmptyDirCreatesTables(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
db, err := OpenWriter(dir, "FULL")
|
|
||||||
|
db, err := Open(dir, "FULL")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
defer func() { _ = db.Close() }()
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
if migErr := Migrate(db); migErr != nil {
|
assertMigrationCount(t, db.Write, 2)
|
||||||
t.Fatal(migErr)
|
for _, table := range []string{
|
||||||
|
"endpoints", "settings", "talk_grants", "groups", "group_members",
|
||||||
|
"messages", "message_bodies", "deliveries", "receipts", "send_keys",
|
||||||
|
"admin_sessions", "api_tokens", "schema_migrations",
|
||||||
|
} {
|
||||||
|
assertTableExists(t, db.Write, table)
|
||||||
}
|
}
|
||||||
if migErr := Migrate(db); migErr != nil {
|
// 空目录首次启动不应产生迁移备份。
|
||||||
t.Fatal(migErr)
|
if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 {
|
||||||
|
t.Fatalf("unexpected backup files on empty start: %d", len(entries))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenIdempotentNoRemigrate(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
db1, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assertMigrationCount(t, db1.Write, 2)
|
||||||
|
_ = db1.Close()
|
||||||
|
|
||||||
|
db2, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db2.Close() }()
|
||||||
|
assertMigrationCount(t, db2.Write, 2)
|
||||||
|
|
||||||
|
if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 {
|
||||||
|
t.Fatalf("idempotent reopen should not backup: %v", entries)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMigrateBackupWhenNewVersion(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
// 模拟仅应用了 0001 的旧库。
|
||||||
|
w, err := OpenWriter(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, execErr := w.Exec(`
|
||||||
|
CREATE TABLE schema_migrations (
|
||||||
|
version INTEGER PRIMARY KEY,
|
||||||
|
applied_at INTEGER NOT NULL
|
||||||
|
)`); execErr != nil {
|
||||||
|
t.Fatal(execErr)
|
||||||
|
}
|
||||||
|
if _, execErr := w.Exec(`INSERT INTO schema_migrations(version, applied_at) VALUES(1, ?)`, time.Now().UnixMilli()); execErr != nil {
|
||||||
|
t.Fatal(execErr)
|
||||||
|
}
|
||||||
|
_ = w.Close()
|
||||||
|
|
||||||
|
db, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
assertMigrationCount(t, db.Write, 2)
|
||||||
|
assertTableExists(t, db.Write, "endpoints")
|
||||||
|
assertTableExists(t, db.Write, "api_tokens")
|
||||||
|
|
||||||
|
backupDir := filepath.Join(dir, "backup")
|
||||||
|
entries, err := os.ReadDir(backupDir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("backup dir missing: %v", err)
|
||||||
|
}
|
||||||
|
found := false
|
||||||
|
for _, e := range entries {
|
||||||
|
if strings.HasPrefix(e.Name(), "pre-migrate-") && strings.HasSuffix(e.Name(), ".db") {
|
||||||
|
found = true
|
||||||
|
info, statErr := e.Info()
|
||||||
|
if statErr != nil {
|
||||||
|
t.Fatal(statErr)
|
||||||
|
}
|
||||||
|
if info.Size() == 0 {
|
||||||
|
t.Fatalf("backup file empty: %s", e.Name())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
t.Fatal("expected pre-migrate-*.db backup")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteQueueOneOpOneTx(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dir := t.TempDir()
|
||||||
|
db, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, execErr := tx.Exec(
|
||||||
|
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
||||||
|
"registration_enabled", "0", time.Now().UnixMilli(),
|
||||||
|
)
|
||||||
|
return execErr
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var n int
|
var value string
|
||||||
if scanErr := db.QueryRow(`SELECT COUNT(*) FROM schema_migrations`).Scan(&n); scanErr != nil {
|
if scanErr := db.Read.QueryRow(`SELECT value FROM settings WHERE key = ?`, "registration_enabled").Scan(&value); scanErr != nil {
|
||||||
t.Fatal(scanErr)
|
t.Fatal(scanErr)
|
||||||
}
|
}
|
||||||
if n != 1 {
|
if value != "0" {
|
||||||
t.Fatalf("want 1 migration row, got %d", n)
|
t.Fatalf("want 0, got %q", value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertMigrationCount(t *testing.T, db *sql.DB, want int) {
|
||||||
|
t.Helper()
|
||||||
|
var n int
|
||||||
|
if err := db.QueryRow(`SELECT COUNT(*) FROM schema_migrations`).Scan(&n); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != want {
|
||||||
|
t.Fatalf("migration rows: want %d got %d", want, n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertTableExists(t *testing.T, db *sql.DB, name string) {
|
||||||
|
t.Helper()
|
||||||
|
var got string
|
||||||
|
err := db.QueryRow(
|
||||||
|
`SELECT name FROM sqlite_master WHERE type='table' AND name=?`,
|
||||||
|
name,
|
||||||
|
).Scan(&got)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("table %s missing: %v", name, err)
|
||||||
}
|
}
|
||||||
nested, openErr := OpenWriter(filepath.Join(dir, "nested"), "NORMAL")
|
|
||||||
if openErr != nil {
|
|
||||||
t.Fatal(openErr)
|
|
||||||
}
|
|
||||||
_ = nested.Close()
|
|
||||||
}
|
}
|
||||||
|
|||||||
+73
-16
@@ -5,6 +5,8 @@ import (
|
|||||||
"embed"
|
"embed"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io/fs"
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"sort"
|
"sort"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -14,8 +16,11 @@ import (
|
|||||||
//go:embed migrations/*.sql
|
//go:embed migrations/*.sql
|
||||||
var migrationFS embed.FS
|
var migrationFS embed.FS
|
||||||
|
|
||||||
// Migrate 应用尚未执行的迁移。T0.1 仅为空执行器加 0001 占位;完整备份与表结构见 T0.3。
|
// Migrate 应用尚未执行的嵌入迁移。
|
||||||
func Migrate(db *sql.DB) error {
|
// 若 dbExisted 为 true(调用 Open/Migrate 前已有 nixmsg.db)且存在未应用版本,
|
||||||
|
// 先 VACUUM INTO <data_dir>/backup/pre-migrate-<时间>.db,再迁移。
|
||||||
|
// 任一版本失败则返回错误,调用方不得继续带半新半旧库提供服务。
|
||||||
|
func Migrate(db *sql.DB, dataDir string, dbExisted bool) error {
|
||||||
if _, err := db.Exec(`
|
if _, err := db.Exec(`
|
||||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||||
version INTEGER PRIMARY KEY,
|
version INTEGER PRIMARY KEY,
|
||||||
@@ -24,9 +29,38 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
|
|||||||
return fmt.Errorf("ensure schema_migrations: %w", err)
|
return fmt.Errorf("ensure schema_migrations: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pending, err := pendingMigrations(db)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if len(pending) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if dbExisted {
|
||||||
|
if err := backupBeforeMigrate(db, dataDir); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, m := range pending {
|
||||||
|
if err := applyMigration(db, m); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type migrationFile struct {
|
||||||
|
version int
|
||||||
|
name string
|
||||||
|
body string
|
||||||
|
}
|
||||||
|
|
||||||
|
func pendingMigrations(db *sql.DB) ([]migrationFile, error) {
|
||||||
entries, err := fs.ReadDir(migrationFS, "migrations")
|
entries, err := fs.ReadDir(migrationFS, "migrations")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("read migrations: %w", err)
|
return nil, fmt.Errorf("read migrations: %w", err)
|
||||||
}
|
}
|
||||||
var names []string
|
var names []string
|
||||||
for _, e := range entries {
|
for _, e := range entries {
|
||||||
@@ -37,10 +71,11 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
|
|||||||
}
|
}
|
||||||
sort.Strings(names)
|
sort.Strings(names)
|
||||||
|
|
||||||
|
var pending []migrationFile
|
||||||
for _, name := range names {
|
for _, name := range names {
|
||||||
version, err := parseMigrationVersion(name)
|
version, err := parseMigrationVersion(name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
var exists int
|
var exists int
|
||||||
err = db.QueryRow(`SELECT 1 FROM schema_migrations WHERE version = ?`, version).Scan(&exists)
|
err = db.QueryRow(`SELECT 1 FROM schema_migrations WHERE version = ?`, version).Scan(&exists)
|
||||||
@@ -48,35 +83,57 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if err != sql.ErrNoRows {
|
if err != sql.ErrNoRows {
|
||||||
return fmt.Errorf("check migration %d: %w", version, err)
|
return nil, fmt.Errorf("check migration %d: %w", version, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
body, err := migrationFS.ReadFile("migrations/" + name)
|
body, err := migrationFS.ReadFile("migrations/" + name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("read migration %s: %w", name, err)
|
return nil, fmt.Errorf("read migration %s: %w", name, err)
|
||||||
}
|
}
|
||||||
sqlText := strings.TrimSpace(string(body))
|
pending = append(pending, migrationFile{
|
||||||
|
version: version,
|
||||||
|
name: name,
|
||||||
|
body: strings.TrimSpace(string(body)),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return pending, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func backupBeforeMigrate(db *sql.DB, dataDir string) error {
|
||||||
|
backupDir := filepath.Join(dataDir, "backup")
|
||||||
|
if err := os.MkdirAll(backupDir, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("mkdir backup: %w", err)
|
||||||
|
}
|
||||||
|
stamp := time.Now().UTC().Format("20060102T150405")
|
||||||
|
backupPath := filepath.Join(backupDir, "pre-migrate-"+stamp+".db")
|
||||||
|
// SQLite VACUUM INTO 需要字面量路径;统一用斜杠,并对单引号转义。
|
||||||
|
quoted := strings.ReplaceAll(filepath.ToSlash(backupPath), "'", "''")
|
||||||
|
if _, err := db.Exec("VACUUM INTO '" + quoted + "'"); err != nil {
|
||||||
|
return fmt.Errorf("vacuum into backup %s: %w", backupPath, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func applyMigration(db *sql.DB, m migrationFile) error {
|
||||||
tx, err := db.Begin()
|
tx, err := db.Begin()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("begin migration %d: %w", version, err)
|
return fmt.Errorf("begin migration %d: %w", m.version, err)
|
||||||
}
|
}
|
||||||
if sqlText != "" {
|
if m.body != "" {
|
||||||
if _, err := tx.Exec(sqlText); err != nil {
|
if _, err := tx.Exec(m.body); err != nil {
|
||||||
_ = tx.Rollback()
|
_ = tx.Rollback()
|
||||||
return fmt.Errorf("apply migration %d: %w", version, err)
|
return fmt.Errorf("apply migration %d (%s): %w", m.version, m.name, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if _, err := tx.Exec(
|
if _, err := tx.Exec(
|
||||||
`INSERT INTO schema_migrations(version, applied_at) VALUES(?, ?)`,
|
`INSERT INTO schema_migrations(version, applied_at) VALUES(?, ?)`,
|
||||||
version,
|
m.version,
|
||||||
time.Now().UnixMilli(),
|
time.Now().UnixMilli(),
|
||||||
); err != nil {
|
); err != nil {
|
||||||
_ = tx.Rollback()
|
_ = tx.Rollback()
|
||||||
return fmt.Errorf("record migration %d: %w", version, err)
|
return fmt.Errorf("record migration %d: %w", m.version, err)
|
||||||
}
|
}
|
||||||
if err := tx.Commit(); err != nil {
|
if err := tx.Commit(); err != nil {
|
||||||
return fmt.Errorf("commit migration %d: %w", version, err)
|
return fmt.Errorf("commit migration %d: %w", m.version, err)
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,131 @@
|
|||||||
|
-- 完整业务表(T0.1 的 0001 仅为 SELECT 1 占位且可能已记入 schema_migrations,故放在 0002)。
|
||||||
|
CREATE TABLE endpoints (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL DEFAULT '',
|
||||||
|
remark TEXT NOT NULL DEFAULT '',
|
||||||
|
source TEXT NOT NULL DEFAULT 'admin',
|
||||||
|
login_hash TEXT NOT NULL,
|
||||||
|
talk_hash TEXT,
|
||||||
|
talk_version INTEGER NOT NULL DEFAULT 0,
|
||||||
|
default_delay_ms INTEGER NOT NULL DEFAULT 0,
|
||||||
|
enabled INTEGER NOT NULL DEFAULT 1,
|
||||||
|
created_at INTEGER NOT NULL,
|
||||||
|
online_since INTEGER,
|
||||||
|
offline_since INTEGER,
|
||||||
|
session_hash TEXT,
|
||||||
|
session_issued_at INTEGER,
|
||||||
|
session_used_at INTEGER
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE settings (
|
||||||
|
key TEXT PRIMARY KEY,
|
||||||
|
value TEXT NOT NULL,
|
||||||
|
updated_at INTEGER NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE talk_grants (
|
||||||
|
sender_id TEXT NOT NULL,
|
||||||
|
target_id TEXT NOT NULL,
|
||||||
|
target_talk_version INTEGER NOT NULL,
|
||||||
|
kind TEXT NOT NULL,
|
||||||
|
created_at INTEGER NOT NULL,
|
||||||
|
PRIMARY KEY (sender_id, target_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE groups (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
owner_id TEXT NOT NULL,
|
||||||
|
created_at INTEGER NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE group_members (
|
||||||
|
group_id TEXT NOT NULL,
|
||||||
|
endpoint_id TEXT NOT NULL,
|
||||||
|
joined_at INTEGER NOT NULL,
|
||||||
|
PRIMARY KEY (group_id, endpoint_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE messages (
|
||||||
|
seq INTEGER PRIMARY KEY,
|
||||||
|
id TEXT NOT NULL,
|
||||||
|
sender_id TEXT NOT NULL,
|
||||||
|
dest_kind TEXT NOT NULL,
|
||||||
|
dest_id TEXT NOT NULL,
|
||||||
|
meta TEXT NOT NULL DEFAULT '{}',
|
||||||
|
content_type TEXT NOT NULL,
|
||||||
|
body_enc TEXT NOT NULL,
|
||||||
|
send_at INTEGER NOT NULL,
|
||||||
|
keep INTEGER NOT NULL,
|
||||||
|
ttl_seconds INTEGER NOT NULL DEFAULT 0,
|
||||||
|
receipt INTEGER NOT NULL,
|
||||||
|
state TEXT NOT NULL,
|
||||||
|
reason TEXT NOT NULL DEFAULT '',
|
||||||
|
created_at INTEGER NOT NULL,
|
||||||
|
UNIQUE (sender_id, id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE message_bodies (
|
||||||
|
seq INTEGER PRIMARY KEY REFERENCES messages(seq) ON DELETE CASCADE,
|
||||||
|
body BLOB NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE deliveries (
|
||||||
|
seq INTEGER NOT NULL REFERENCES messages(seq) ON DELETE CASCADE,
|
||||||
|
endpoint_id TEXT NOT NULL,
|
||||||
|
send_at INTEGER NOT NULL,
|
||||||
|
keep INTEGER NOT NULL,
|
||||||
|
state TEXT NOT NULL,
|
||||||
|
reason TEXT NOT NULL DEFAULT '',
|
||||||
|
expire_at INTEGER,
|
||||||
|
pushed_conn TEXT,
|
||||||
|
pushed_at INTEGER,
|
||||||
|
attempts INTEGER NOT NULL DEFAULT 0,
|
||||||
|
updated_at INTEGER NOT NULL,
|
||||||
|
PRIMARY KEY (seq, endpoint_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE receipts (
|
||||||
|
receipt_id INTEGER PRIMARY KEY,
|
||||||
|
sender_id TEXT NOT NULL,
|
||||||
|
msg_id TEXT NOT NULL,
|
||||||
|
endpoint_id TEXT NOT NULL DEFAULT '',
|
||||||
|
state TEXT NOT NULL,
|
||||||
|
reason TEXT NOT NULL DEFAULT '',
|
||||||
|
created_at INTEGER NOT NULL,
|
||||||
|
acked INTEGER NOT NULL DEFAULT 0
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE send_keys (
|
||||||
|
sender_id TEXT NOT NULL,
|
||||||
|
msg_id TEXT NOT NULL,
|
||||||
|
request_sha256 BLOB NOT NULL,
|
||||||
|
created_at INTEGER NOT NULL,
|
||||||
|
PRIMARY KEY (sender_id, msg_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE admin_sessions (
|
||||||
|
token_hash TEXT PRIMARY KEY,
|
||||||
|
created_at INTEGER NOT NULL,
|
||||||
|
expires_at INTEGER NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE api_tokens (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
token_hash TEXT NOT NULL UNIQUE,
|
||||||
|
enabled INTEGER NOT NULL DEFAULT 1,
|
||||||
|
created_at INTEGER NOT NULL,
|
||||||
|
last_used_at INTEGER
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX idx_messages_due ON messages(state, send_at);
|
||||||
|
CREATE INDEX idx_messages_dest ON messages(dest_kind, dest_id, state);
|
||||||
|
CREATE INDEX idx_messages_created ON messages(created_at);
|
||||||
|
CREATE INDEX idx_messages_sender ON messages(sender_id, created_at);
|
||||||
|
CREATE INDEX idx_messages_sender_state ON messages(sender_id, state);
|
||||||
|
CREATE INDEX idx_deliveries_outbox ON deliveries(endpoint_id, state, send_at, seq);
|
||||||
|
CREATE INDEX idx_deliveries_expire ON deliveries(state, expire_at);
|
||||||
|
CREATE INDEX idx_receipts_outbox ON receipts(sender_id, acked, receipt_id);
|
||||||
|
CREATE INDEX idx_talk_grants_target ON talk_grants(target_id);
|
||||||
|
CREATE INDEX idx_group_members_endpoint ON group_members(endpoint_id);
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrQueueClosed 表示写入队列已关闭。
|
||||||
|
var ErrQueueClosed = errors.New("store: write queue closed")
|
||||||
|
|
||||||
|
// WriteFunc 在单个写事务中执行的操作。
|
||||||
|
type WriteFunc func(tx *sql.Tx) error
|
||||||
|
|
||||||
|
// Queue 写入队列:提交一个写操作并拿到结果。
|
||||||
|
//
|
||||||
|
// 本任务(T0.3)实现为互斥串行的一操作一事务,不做合并;DEVELOPMENT 7.2
|
||||||
|
// 要求的写 goroutine 合并提交(最多 256 个或凑满 2ms、SAVEPOINT 隔离失败)留给 P2。
|
||||||
|
type Queue struct {
|
||||||
|
db *sql.DB
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
closed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewQueue 创建简单写入队列(一操作一事务)。
|
||||||
|
func NewQueue(db *sql.DB) *Queue {
|
||||||
|
return &Queue{db: db}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Do 提交写操作并等待提交结果。
|
||||||
|
func (q *Queue) Do(ctx context.Context, fn WriteFunc) error {
|
||||||
|
if fn == nil {
|
||||||
|
return errors.New("store: nil write func")
|
||||||
|
}
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
if q.closed {
|
||||||
|
return ErrQueueClosed
|
||||||
|
}
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
tx, err := q.db.BeginTx(ctx, nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := fn(tx); err != nil {
|
||||||
|
_ = tx.Rollback()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return tx.Commit()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close 关闭队列,之后 Do 返回 ErrQueueClosed。
|
||||||
|
func (q *Queue) Close() error {
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
q.closed = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user