package identity import ( "context" "database/sql" "errors" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) // SelfGet 返回自己的资料。 func (a *App) SelfGet(ctx context.Context, endpointID string) (SelfInfo, error) { if !protocol.ValidEndpointID(endpointID) { return SelfInfo{}, errCode(protocol.CodeBadRequest, "invalid endpoint id") } var info SelfInfo var talk sql.NullString err := a.db.Read.QueryRowContext(ctx, ` SELECT id, name, default_delay_ms, talk_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&info.ID, &info.Name, &info.DefaultDelayMs, &talk) if errors.Is(err, sql.ErrNoRows) { return SelfInfo{}, errCode(protocol.CodeNotFound, "endpoint not found") } if err != nil { return SelfInfo{}, err } info.TalkPasswordSet = talk.Valid && talk.String != "" return info, nil } // SelfUpdate 更新名称与默认延迟。 func (a *App) SelfUpdate(ctx context.Context, endpointID string, req *protocol.SelfUpdate) error { if !protocol.ValidEndpointID(endpointID) { return errCode(protocol.CodeBadRequest, "invalid endpoint id") } if req == nil { return errCode(protocol.CodeBadRequest, "nil request") } if err := req.Validate(); err != nil { return err } if req.Name == "" && req.DefaultDelayMs == nil { return errCode(protocol.CodeBadRequest, "nothing to update") } if req.DefaultDelayMs != nil && a.maxScheduleSeconds > 0 { maxMs := a.maxScheduleSeconds * 1000 if *req.DefaultDelayMs > maxMs { return errCode(protocol.CodeBadRequest, "default_delay_ms exceeds max_schedule_seconds") } } return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { var exists int if err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, endpointID).Scan(&exists); err != nil { if errors.Is(err, sql.ErrNoRows) { return errCode(protocol.CodeNotFound, "endpoint not found") } return err } if req.Name != "" { if _, err := tx.Exec(`UPDATE endpoints SET name = ? WHERE id = ?`, req.Name, endpointID); err != nil { return err } } if req.DefaultDelayMs != nil { if _, err := tx.Exec(`UPDATE endpoints SET default_delay_ms = ? WHERE id = ?`, *req.DefaultDelayMs, endpointID); err != nil { return err } } return nil }) } // SelfSetTalkPassword 设置或清除对话密码,并增加 talk_version;改密清零对方锁定计数。 func (a *App) SelfSetTalkPassword(ctx context.Context, endpointID, talkPassword string) error { if !protocol.ValidEndpointID(endpointID) { return errCode(protocol.CodeBadRequest, "invalid endpoint id") } if !protocol.ValidTalkPassword(talkPassword) { return errCode(protocol.CodeBadRequest, "invalid talk_password") } var talkHash any if talkPassword != "" { if a.hash == nil { return errors.New("identity: hash pool required") } phc, err := a.hash.Hash(ctx, auth.PasswordTalk, talkPassword) if err != nil { return err } talkHash = phc } err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { res, err := tx.Exec(` UPDATE endpoints SET talk_hash = ?, talk_version = talk_version + 1 WHERE id = ?`, talkHash, endpointID) if err != nil { return err } n, _ := res.RowsAffected() if n == 0 { return errCode(protocol.CodeNotFound, "endpoint not found") } return nil }) if err != nil { return err } // D24:改密清零按对方计的对话密码失败计数。 a.locks.Clear(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: endpointID}) return nil } // SelfChangeLoginPassword 要求旧密码;成功返回新 session_token,旧令牌作废,当前连接由调用方保留。 func (a *App) SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (string, error) { if !protocol.ValidEndpointID(endpointID) { return "", errCode(protocol.CodeBadRequest, "invalid endpoint id") } if protocol.LoginPasswordForbiddenPrefix(newPassword) { return "", errCode(protocol.CodeBadRequest, "login password must not start with nst_") } if !protocol.ValidLoginPassword(newPassword) || newPassword == "" { return "", errCode(protocol.CodeBadRequest, "invalid new_password") } if a.hash == nil || a.sessions == nil { return "", errors.New("identity: hash/sessions required") } var loginHash string err := a.db.Read.QueryRowContext(ctx, `SELECT login_hash FROM endpoints WHERE id = ?`, endpointID).Scan(&loginHash) if errors.Is(err, sql.ErrNoRows) { return "", errCode(protocol.CodeNotFound, "endpoint not found") } if err != nil { return "", err } if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}); locked { return "", errCode(protocol.CodeRateLimited, "login locked") } if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}); locked { return "", errCode(protocol.CodeRateLimited, "login locked") } ok, err := a.hash.Verify(ctx, auth.PasswordLogin, oldPassword, loginHash) if err != nil { return "", err } if !ok { a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpointIP, EndpointID: endpointID, IP: remoteIP}) a.locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: endpointID}) return "", errCode(protocol.CodeUnauthorized, "old password invalid") } newHash, err := a.hash.Hash(ctx, auth.PasswordLogin, newPassword) if err != nil { return "", err } token, tokenHash, err := a.sessions.Issue(ctx) if err != nil { return "", err } nowMs := a.now().UnixMilli() hashHex := encodeSessionHash(tokenHash) err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { res, e := tx.Exec(` UPDATE endpoints SET login_hash = ?, session_hash = ?, session_issued_at = ?, session_used_at = ? WHERE id = ?`, newHash, hashHex, nowMs, nowMs, endpointID) if e != nil { return e } n, _ := res.RowsAffected() if n == 0 { return errCode(protocol.CodeNotFound, "endpoint not found") } return nil }) if err != nil { return "", err } return token, nil } // SelfLogout 清空会话令牌;若注入了 ConnControl 则断开当前连接。 func (a *App) SelfLogout(ctx context.Context, endpointID string) error { if !protocol.ValidEndpointID(endpointID) { return errCode(protocol.CodeBadRequest, "invalid endpoint id") } err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { res, e := tx.Exec(` UPDATE endpoints SET session_hash = NULL, session_issued_at = NULL, session_used_at = NULL WHERE id = ?`, endpointID) if e != nil { return e } n, _ := res.RowsAffected() if n == 0 { return errCode(protocol.CodeNotFound, "endpoint not found") } return nil }) if err != nil { return err } if a.connCtrl != nil { _ = a.connCtrl.Disconnect(ctx, endpointID, "", port.DisconnectNormal) } return nil } // Disable / Enable / Delete 属 I5,此处保留未实现。 func (a *App) Disable(context.Context, string) error { return ErrNotImplemented } func (a *App) Enable(context.Context, string) error { return ErrNotImplemented } func (a *App) Delete(context.Context, string) error { return ErrNotImplemented }