feat: sync AI compose and ChatBot features

This commit is contained in:
A 2026-07-03 19:43:20 +08:00
parent 35e5d38f4d
commit b7269b135f
75 changed files with 5426 additions and 123 deletions

View file

@ -2,15 +2,344 @@ package rpc
import (
"context"
"errors"
"strings"
"github.com/gotd/td/tg"
"github.com/gotd/td/tgerr"
"telesrv/internal/compat/tdesktop"
"telesrv/internal/domain"
)
// registerAiCompose 注册第一阶段 TDesktop 启动所需 aicompose.* RPC 兼容响应。
func (r *Router) registerAiCompose(d *tg.ServerDispatcher) {
d.OnAicomposeGetTones(func(ctx context.Context, hash int64) (tg.AicomposeTonesClass, error) {
return tdesktop.AiComposeTones(), nil
d.OnAicomposeGetTones(r.onAicomposeGetTones)
d.OnAicomposeCreateTone(r.onAicomposeCreateTone)
d.OnAicomposeUpdateTone(r.onAicomposeUpdateTone)
d.OnAicomposeSaveTone(r.onAicomposeSaveTone)
d.OnAicomposeDeleteTone(r.onAicomposeDeleteTone)
d.OnAicomposeGetTone(r.onAicomposeGetTone)
d.OnAicomposeGetToneExample(r.onAicomposeGetToneExample)
}
func (r *Router) onAicomposeGetTones(ctx context.Context, hash int64) (tg.AicomposeTonesClass, error) {
if r.deps.AICompose == nil {
return &tg.AicomposeTones{Hash: 0, Tones: []tg.AiComposeToneClass{}, Users: []tg.UserClass{}}, nil
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return nil, err
}
tones, notModified, err := r.deps.AICompose.ListTones(ctx, userID, hash)
if err != nil {
return nil, aiComposeErr(err)
}
if notModified {
return &tg.AicomposeTonesNotModified{}, nil
}
return r.tgAIComposeTones(ctx, userID, tones), nil
}
func (r *Router) onAicomposeCreateTone(ctx context.Context, req *tg.AicomposeCreateToneRequest) (tg.AiComposeToneClass, error) {
if r.deps.AICompose == nil {
return nil, tgerr.New(400, "AICOMPOSE_TONE_INVALID")
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return nil, err
}
tone, err := r.deps.AICompose.CreateTone(ctx, domain.AIComposeToneInput{
UserID: userID,
DisplayAuthor: req.DisplayAuthor,
EmojiID: req.EmojiID,
Title: req.Title,
Prompt: req.Prompt,
})
if err != nil {
return nil, aiComposeErr(err)
}
r.pushAIComposeTonesChanged(ctx, userID)
return r.tgAIComposeTone(ctx, userID, tone), nil
}
func (r *Router) onAicomposeUpdateTone(ctx context.Context, req *tg.AicomposeUpdateToneRequest) (tg.AiComposeToneClass, error) {
if r.deps.AICompose == nil {
return nil, tgerr.New(400, "AICOMPOSE_TONE_INVALID")
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return nil, err
}
ref, err := domainAIComposeToneRef(req.Tone)
if err != nil {
return nil, err
}
update := domain.AIComposeToneUpdate{
Ref: ref,
UserID: userID,
}
if req.Flags.Has(0) {
v := req.DisplayAuthor
update.DisplayAuthor = &v
}
if req.Flags.Has(1) {
v := req.EmojiID
update.EmojiID = &v
}
if req.Flags.Has(2) {
v := req.Title
update.Title = &v
}
if req.Flags.Has(3) {
v := req.Prompt
update.Prompt = &v
}
tone, err := r.deps.AICompose.UpdateTone(ctx, update)
if err != nil {
return nil, aiComposeErr(err)
}
r.pushAIComposeTonesChanged(ctx, userID)
return r.tgAIComposeTone(ctx, userID, tone), nil
}
func (r *Router) onAicomposeSaveTone(ctx context.Context, req *tg.AicomposeSaveToneRequest) (bool, error) {
if r.deps.AICompose == nil {
return false, tgerr.New(400, "AICOMPOSE_TONE_INVALID")
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return false, err
}
ref, err := domainAIComposeToneRef(req.Tone)
if err != nil {
return false, err
}
if err := r.deps.AICompose.SaveTone(ctx, userID, ref, req.Unsave); err != nil {
return false, aiComposeErr(err)
}
r.pushAIComposeTonesChanged(ctx, userID)
return true, nil
}
func (r *Router) onAicomposeDeleteTone(ctx context.Context, tone tg.InputAiComposeToneClass) (bool, error) {
if r.deps.AICompose == nil {
return false, tgerr.New(400, "AICOMPOSE_TONE_INVALID")
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return false, err
}
ref, err := domainAIComposeToneRef(tone)
if err != nil {
return false, err
}
if err := r.deps.AICompose.DeleteTone(ctx, userID, ref); err != nil {
return false, aiComposeErr(err)
}
r.pushAIComposeTonesChanged(ctx, userID)
return true, nil
}
func (r *Router) onAicomposeGetTone(ctx context.Context, tone tg.InputAiComposeToneClass) (tg.AicomposeTonesClass, error) {
if r.deps.AICompose == nil {
return &tg.AicomposeTones{Hash: 0, Tones: []tg.AiComposeToneClass{}, Users: []tg.UserClass{}}, nil
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return nil, err
}
ref, err := domainAIComposeToneRef(tone)
if err != nil {
return nil, err
}
tones, err := r.deps.AICompose.GetTone(ctx, userID, ref)
if err != nil {
return nil, aiComposeErr(err)
}
return r.tgAIComposeTones(ctx, userID, tones), nil
}
func (r *Router) onAicomposeGetToneExample(ctx context.Context, req *tg.AicomposeGetToneExampleRequest) (*tg.AiComposeToneExample, error) {
if r.deps.AICompose == nil {
return nil, tgerr.New(400, "AICOMPOSE_TONE_INVALID")
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return nil, err
}
ref, err := domainAIComposeToneRef(req.Tone)
if err != nil {
return nil, err
}
example, err := r.deps.AICompose.GetToneExample(ctx, userID, ref, req.Num)
if err != nil {
return nil, aiComposeErr(err)
}
return &tg.AiComposeToneExample{
From: tgAIComposeText(example.From),
To: tgAIComposeText(example.To),
}, nil
}
func (r *Router) onMessagesComposeMessageWithAI(ctx context.Context, req *tg.MessagesComposeMessageWithAIRequest) (*tg.MessagesComposedMessageWithAI, error) {
if r.deps.AICompose == nil {
return nil, tgerr.New(500, "AICOMPOSE_FAILED")
}
userID, err := r.currentAIComposeUserID(ctx)
if err != nil {
return nil, err
}
ref := domain.AIComposeToneRef{}
if req.Flags.Has(2) || req.Tone != nil {
ref, err = domainAIComposeOptionalToneRef(req.Tone)
if err != nil {
return nil, err
}
}
in := domain.AIComposeRequest{
UserID: userID,
Text: domainAIComposeText(userID, req.Text),
Proofread: req.Proofread,
Emojify: req.Emojify,
TranslateToLang: req.TranslateToLang,
Tone: ref,
}
result, err := r.deps.AICompose.Compose(ctx, in)
if err != nil {
return nil, aiComposeErr(err)
}
out := &tg.MessagesComposedMessageWithAI{
ResultText: tgAIComposeText(result.ResultText),
}
if result.DiffText != nil {
out.SetDiffText(tgAIComposeText(*result.DiffText))
}
return out, nil
}
func (r *Router) currentAIComposeUserID(ctx context.Context) (int64, error) {
userID, ok, err := r.currentUserID(ctx)
if err != nil {
return 0, internalErr()
}
if !ok {
return 0, authKeyUnregisteredErr()
}
return userID, nil
}
func domainAIComposeToneRef(tone tg.InputAiComposeToneClass) (domain.AIComposeToneRef, error) {
switch t := tone.(type) {
case *tg.InputAiComposeToneDefault:
return domain.AIComposeToneRef{Kind: domain.AIComposeToneRefDefault, DefaultTone: t.Tone}, nil
case *tg.InputAiComposeToneID:
return domain.AIComposeToneRef{Kind: domain.AIComposeToneRefID, ID: t.ID, AccessHash: t.AccessHash}, nil
case *tg.InputAiComposeToneSlug:
return domain.AIComposeToneRef{Kind: domain.AIComposeToneRefSlug, Slug: t.Slug}, nil
default:
return domain.AIComposeToneRef{}, inputConstructorInvalidErr()
}
}
func domainAIComposeOptionalToneRef(tone tg.InputAiComposeToneClass) (domain.AIComposeToneRef, error) {
switch t := tone.(type) {
case *tg.InputAiComposeToneDefault:
if strings.TrimSpace(t.Tone) == "" {
return domain.AIComposeToneRef{}, nil
}
}
return domainAIComposeToneRef(tone)
}
func domainAIComposeText(viewerUserID int64, in tg.TextWithEntities) domain.AIComposeText {
return domain.AIComposeText{
Text: in.Text,
Entities: domainMessageEntitiesForViewer(viewerUserID, in.Entities),
}
}
func tgAIComposeText(in domain.AIComposeText) tg.TextWithEntities {
return tg.TextWithEntities{
Text: in.Text,
Entities: tgMessageEntities(in.Entities),
}
}
func (r *Router) tgAIComposeTones(ctx context.Context, userID int64, in domain.AIComposeTones) tg.AicomposeTonesClass {
tones := make([]tg.AiComposeToneClass, 0, len(in.Tones))
for _, tone := range in.Tones {
tones = append(tones, r.tgAIComposeTone(ctx, userID, tone))
}
return &tg.AicomposeTones{
Hash: in.Hash,
Tones: tones,
Users: []tg.UserClass{},
}
}
func (r *Router) tgAIComposeTone(_ context.Context, userID int64, in domain.AIComposeTone) tg.AiComposeToneClass {
if in.Default {
return &tg.AiComposeToneDefault{
Tone: in.Slug,
EmojiID: in.EmojiID,
Title: in.Title,
}
}
out := &tg.AiComposeTone{
ID: in.ID,
AccessHash: in.AccessHash,
Slug: in.Slug,
Title: in.Title,
}
out.SetCreator(in.Creator || in.OwnerUserID == userID)
if in.EmojiID != 0 {
out.SetEmojiID(in.EmojiID)
}
if in.Prompt != "" {
out.SetPrompt(in.Prompt)
}
if in.InstallsCount > 0 {
out.SetInstallsCount(in.InstallsCount)
}
if in.AuthorID != 0 {
out.SetAuthorID(in.AuthorID)
}
if in.ExampleEnglish != nil {
out.SetExampleEnglish(tg.AiComposeToneExample{
From: tgAIComposeText(in.ExampleEnglish.From),
To: tgAIComposeText(in.ExampleEnglish.To),
})
}
return out
}
func (r *Router) pushAIComposeTonesChanged(ctx context.Context, userID int64) {
r.pushUserMessageTransient(ctx, userID, "push ai compose tones update", &tg.UpdateShort{
Update: &tg.UpdateAiComposeTones{},
Date: int(r.clock.Now().Unix()),
})
}
func aiComposeErr(err error) error {
switch {
case err == nil:
return nil
case errors.Is(err, domain.ErrAIComposeToneNotFound):
return tgerr.New(400, "TONE_NOT_FOUND")
case errors.Is(err, domain.ErrAIComposeToneInvalid):
return tgerr.New(400, "AICOMPOSE_TONE_INVALID")
case errors.Is(err, domain.ErrAIComposeToneLimitExceeded):
return tgerr.New(400, "TONES_SAVED_TOO_MANY")
case errors.Is(err, domain.ErrAIComposeRateLimited):
return floodWaitErr(60)
case errors.Is(err, domain.ErrAIComposeInvalid):
return inputRequestInvalidErr()
case errors.Is(err, domain.ErrAIComposeDisabled):
return tgerr.New(400, "AICOMPOSE_DISABLED")
case errors.Is(err, domain.ErrAIComposeProviderTimeout):
return tgerr.New(500, "AICOMPOSE_TIMEOUT")
case errors.Is(err, domain.ErrAIComposeProviderUnavailable):
return tgerr.New(500, "AICOMPOSE_FAILED")
default:
return internalErr()
}
}

View file

@ -0,0 +1,191 @@
package rpc
import (
"context"
"testing"
"github.com/gotd/td/clock"
"github.com/gotd/td/tg"
"github.com/gotd/td/tgerr"
"go.uber.org/zap/zaptest"
aiapp "telesrv/internal/app/ai"
"telesrv/internal/domain"
"telesrv/internal/store/memory"
)
func newAIComposeTestRouter(t *testing.T) *Router {
t.Helper()
return New(Config{}, Deps{
AICompose: aiapp.NewService(memory.NewAIComposeStore()),
}, zaptest.NewLogger(t), clock.System)
}
func TestAIComposeGetTonesReturnsDefaultsAndNotModified(t *testing.T) {
r := newAIComposeTestRouter(t)
ctx := WithUserID(context.Background(), 1001)
got, err := r.onAicomposeGetTones(ctx, 0)
if err != nil {
t.Fatalf("getTones: %v", err)
}
tones, ok := got.(*tg.AicomposeTones)
if !ok {
t.Fatalf("getTones = %T, want *tg.AicomposeTones", got)
}
if tones.Hash == 0 || len(tones.Tones) == 0 {
t.Fatalf("getTones hash/tones = %d/%d, want non-empty", tones.Hash, len(tones.Tones))
}
again, err := r.onAicomposeGetTones(ctx, tones.Hash)
if err != nil {
t.Fatalf("getTones(hash): %v", err)
}
if _, ok := again.(*tg.AicomposeTonesNotModified); !ok {
t.Fatalf("getTones(hash) = %T, want tonesNotModified", again)
}
}
func TestMessagesComposeMessageWithAIProofread(t *testing.T) {
r := newAIComposeTestRouter(t)
ctx := WithUserID(context.Background(), 1001)
got, err := r.onMessagesComposeMessageWithAI(ctx, &tg.MessagesComposeMessageWithAIRequest{
Proofread: true,
Text: tg.TextWithEntities{Text: "hello world"},
})
if err != nil {
t.Fatalf("composeMessageWithAI: %v", err)
}
if got.ResultText.Text != "hello world." {
t.Fatalf("compose result = %q, want local polished text", got.ResultText.Text)
}
diff, ok := got.GetDiffText()
if !ok {
t.Fatal("DiffText = nil, want proofread diff")
}
if len(diff.Entities) != 1 {
t.Fatalf("DiffText entities = %d, want 1", len(diff.Entities))
}
replace, ok := diff.Entities[0].(*tg.MessageEntityDiffReplace)
if !ok {
t.Fatalf("DiffText entity = %T, want *tg.MessageEntityDiffReplace", diff.Entities[0])
}
if replace.OldText != "hello world" || replace.Offset != 0 || replace.Length != len([]rune(got.ResultText.Text)) {
t.Fatalf("DiffText replace = %#v", replace)
}
}
func TestMessagesComposeMessageWithAIEmptyDefaultToneIsOptional(t *testing.T) {
r := newAIComposeTestRouter(t)
ctx := WithUserID(context.Background(), 1001)
req := &tg.MessagesComposeMessageWithAIRequest{
Text: tg.TextWithEntities{Text: "hello world"},
}
req.SetTranslateToLang("en")
req.SetTone(&tg.InputAiComposeToneDefault{})
got, err := r.onMessagesComposeMessageWithAI(ctx, req)
if err != nil {
t.Fatalf("composeMessageWithAI empty default tone: %v", err)
}
if got.ResultText.Text == "" {
t.Fatal("compose result is empty")
}
if _, err := r.onAicomposeGetTone(ctx, &tg.InputAiComposeToneDefault{}); !tgerr.Is(err, "TONE_NOT_FOUND") {
t.Fatalf("getTone empty default err = %v, want TONE_NOT_FOUND", err)
}
}
func TestAIComposeCustomToneCRUD(t *testing.T) {
r := newAIComposeTestRouter(t)
ctx := WithUserID(context.Background(), 1001)
createdRaw, err := r.onAicomposeCreateTone(ctx, &tg.AicomposeCreateToneRequest{
DisplayAuthor: true,
Title: "Sharp",
Prompt: "Make it crisp.",
})
if err != nil {
t.Fatalf("createTone: %v", err)
}
created, ok := createdRaw.(*tg.AiComposeTone)
if !ok {
t.Fatalf("createTone = %T, want *tg.AiComposeTone", createdRaw)
}
if created.ID == 0 || created.AccessHash == 0 || created.Slug == "" || !created.GetCreator() {
t.Fatalf("created tone = %#v", created)
}
update := &tg.AicomposeUpdateToneRequest{Tone: &tg.InputAiComposeToneID{ID: created.ID, AccessHash: created.AccessHash}}
update.SetTitle("Brief")
updatedRaw, err := r.onAicomposeUpdateTone(ctx, update)
if err != nil {
t.Fatalf("updateTone: %v", err)
}
updated := updatedRaw.(*tg.AiComposeTone)
if updated.Title != "Brief" {
t.Fatalf("updated title = %q, want Brief", updated.Title)
}
got, err := r.onAicomposeGetTone(ctx, &tg.InputAiComposeToneSlug{Slug: created.Slug})
if err != nil {
t.Fatalf("getTone: %v", err)
}
one := got.(*tg.AicomposeTones)
if len(one.Tones) != 1 {
t.Fatalf("getTone tones = %d, want 1", len(one.Tones))
}
if ok, err := r.onAicomposeDeleteTone(ctx, &tg.InputAiComposeToneID{ID: created.ID, AccessHash: created.AccessHash}); err != nil || !ok {
t.Fatalf("deleteTone = %v/%v, want true/nil", ok, err)
}
if _, err := r.onAicomposeGetTone(ctx, &tg.InputAiComposeToneSlug{Slug: created.Slug}); !tgerr.Is(err, "TONE_NOT_FOUND") {
t.Fatalf("getTone after delete err = %v, want TONE_NOT_FOUND", err)
}
}
func TestAIComposeToneLimitUsesClientError(t *testing.T) {
r := newAIComposeTestRouter(t)
ctx := WithUserID(context.Background(), 1001)
for i := 0; i < domain.AIComposeToneSavedLimitDefault; i++ {
if _, err := r.onAicomposeCreateTone(ctx, &tg.AicomposeCreateToneRequest{
Title: "Tone",
Prompt: "Make it crisp.",
}); err != nil {
t.Fatalf("createTone %d: %v", i, err)
}
}
_, err := r.onAicomposeCreateTone(ctx, &tg.AicomposeCreateToneRequest{
Title: "Extra",
Prompt: "Make it crisp.",
})
if !tgerr.Is(err, "TONES_SAVED_TOO_MANY") {
t.Fatalf("createTone over limit err = %v, want TONES_SAVED_TOO_MANY", err)
}
}
func TestAIComposeToneMutationPushesRefreshUpdate(t *testing.T) {
sessions := &captureSessions{}
r := New(Config{}, Deps{
AICompose: aiapp.NewService(memory.NewAIComposeStore()),
Sessions: sessions,
}, zaptest.NewLogger(t), clock.System)
ctx := WithSessionID(WithAuthKeyID(WithUserID(context.Background(), 1001), [8]byte{1}), 77)
if _, err := r.onAicomposeCreateTone(ctx, &tg.AicomposeCreateToneRequest{
Title: "Sharp",
Prompt: "Make it crisp.",
}); err != nil {
t.Fatalf("createTone: %v", err)
}
got := sessions.lastUserPush()
short, ok := got.(*tg.UpdateShort)
if !ok {
t.Fatalf("pushed update = %T, want *tg.UpdateShort", got)
}
if _, ok := short.Update.(*tg.UpdateAiComposeTones); !ok {
t.Fatalf("pushed short update = %T, want *tg.UpdateAiComposeTones", short.Update)
}
if snap := sessions.snapshot(); snap.sessionID != 77 {
t.Fatalf("excluded session = %d, want 77", snap.sessionID)
}
}

View file

@ -0,0 +1,112 @@
package rpc
import (
"context"
"fmt"
"hash/fnv"
"net/url"
"strings"
"time"
"telesrv/internal/domain"
)
const aiComposeToneWebPageType = "telegram_aicomposetone"
func (r *Router) resolveAIComposeStyleWebPage(ctx context.Context, rawURL string) (domain.MessageWebPage, bool) {
link, ok := parseAIComposeStyleLink(rawURL)
if !ok || r.deps.AICompose == nil {
return domain.MessageWebPage{}, false
}
userID, ok, err := r.currentUserID(ctx)
if err != nil || !ok {
return domain.MessageWebPage{}, false
}
tones, err := r.deps.AICompose.GetTone(ctx, userID, domain.AIComposeToneRef{
Kind: domain.AIComposeToneRefSlug,
Slug: link.slug,
})
if err != nil || len(tones.Tones) == 0 {
return domain.MessageWebPage{}, false
}
tone := tones.Tones[0]
now := time.Now()
if r.clock != nil {
now = r.clock.Now()
}
page := domain.MessageWebPage{
State: domain.MessageWebPageStateDone,
ID: domain.WebPageURLHash(link.normalized),
URL: link.normalized,
DisplayURL: link.display,
Hash: aiComposeToneWebPageHash(tone),
Date: int(now.Unix()),
Type: aiComposeToneWebPageType,
SiteName: "Telegram",
Title: tone.Title,
Description: tone.Prompt,
ComposeToneEmojiID: tone.EmojiID,
}
if page.Title == "" {
page.Title = tone.Slug
}
if page.Description == "" {
page.Description = "AI compose style"
}
return page, true
}
type aiComposeStyleLink struct {
normalized string
display string
slug string
}
func parseAIComposeStyleLink(raw string) (aiComposeStyleLink, bool) {
normalized, ok := domain.NormalizeWebPageURL(raw)
if !ok {
return aiComposeStyleLink{}, false
}
u, err := url.Parse(normalized)
if err != nil {
return aiComposeStyleLink{}, false
}
host := strings.ToLower(u.Hostname())
if !aiComposeStyleHostAllowed(host) {
return aiComposeStyleLink{}, false
}
parts := strings.Split(strings.Trim(strings.ToLower(u.EscapedPath()), "/"), "/")
var slug string
switch {
case len(parts) == 1 && parts[0] == "addstyle":
slug = u.Query().Get("slug")
case len(parts) == 2 && parts[0] == "addstyle":
if decoded, err := url.PathUnescape(parts[1]); err == nil {
slug = decoded
}
}
slug = strings.ToLower(strings.TrimSpace(slug))
if slug == "" {
return aiComposeStyleLink{}, false
}
display := host
if path := strings.Trim(u.EscapedPath(), "/"); path != "" {
display += "/" + path
}
return aiComposeStyleLink{normalized: normalized, display: display, slug: slug}, true
}
func aiComposeStyleHostAllowed(host string) bool {
switch host {
case "t.me", "telegram.me", "telesrv.net", "localhost", "127.0.0.1":
return true
default:
return false
}
}
func aiComposeToneWebPageHash(tone domain.AIComposeTone) int {
h := fnv.New32a()
_, _ = fmt.Fprintf(h, "%s|%s|%d|%s|%d", tone.Slug, tone.Title, tone.EmojiID, tone.Prompt, tone.UpdatedAt)
return int(h.Sum32() & 0x7fffffff)
}

View file

@ -10,10 +10,10 @@ import (
"telesrv/internal/domain"
)
// 本文件实现 app/bots 的 RouterHooks 回调:token revoke 后的 session 失效闭环,
// 命令变更后的 updateBotCommands 在线推送,以及 @Stickers 发布后的
// updateStickerSets 在线提示。Router 创建后经
// botsService.SetRouterHooks(router) 装配(见 cmd/telesrv/main.go)。
// 本文件实现 app/bots 的 rpc 回调:token revoke 后的 session 失效闭环,
// 命令变更后的 updateBotCommands 在线推送、@Stickers 发布后的 updateStickerSets
// 在线提示,以及 @ChatBot 流式草稿 transient 推送。Router 创建后经
// botsService.SetRouterHooks / SetTextDraftPusher 装配(见 cmd/telesrv/main.go)。
// maxBotCommandsPushPeers 限制单次命令变更的推送扇出(bot 的最近 dialog peer 数)。
// 超出的离线/长尾用户靠 bot_info_version bump 在下次 getFullUser 时拿到新命令。
@ -81,6 +81,25 @@ func (r *Router) PushStickerSetsChanged(ctx context.Context, userID int64, kind
}()
}
// PushBotTextDraft 推送内置 service bot 的流式文本草稿。草稿是 TDesktop 专用的
// transient typing action,不写 message/dialog/pts/outbox;最终可恢复事实仍由随后
// 入库的普通 bot message 承担。
func (r *Router) PushBotTextDraft(ctx context.Context, botUserID, userID, randomID int64, text string) {
if botUserID == 0 || userID == 0 || randomID == 0 || text == "" {
return
}
r.pushUserMessageTransient(context.WithoutCancel(ctx), userID, "push bot text draft", &tg.UpdateShort{
Update: &tg.UpdateUserTyping{
UserID: botUserID,
Action: &tg.SendMessageTextDraftAction{
RandomID: randomID,
Text: tg.TextWithEntities{Text: text},
},
},
Date: int(r.clock.Now().Unix()),
})
}
func (r *Router) pushBotCommandsChanged(ctx context.Context, botUserID int64, commands []domain.BotCommand) {
defer func() {
if rec := recover(); rec != nil {

View file

@ -0,0 +1,42 @@
package rpc
import (
"context"
"testing"
"github.com/gotd/td/clock"
"github.com/gotd/td/tg"
"go.uber.org/zap/zaptest"
"telesrv/internal/domain"
)
func TestPushBotTextDraftUsesUserTypingDraftAction(t *testing.T) {
sessions := &captureSessions{}
r := New(Config{}, Deps{Sessions: sessions}, zaptest.NewLogger(t), clock.System)
ctx := WithSessionID(WithAuthKeyID(context.Background(), [8]byte{1}), 77)
r.PushBotTextDraft(ctx, domain.ChatBotUserID, 1001, 4242, "Hello from AI")
short, ok := sessions.lastUserPush().(*tg.UpdateShort)
if !ok {
t.Fatalf("pushed update = %T, want *tg.UpdateShort", sessions.lastUserPush())
}
update, ok := short.Update.(*tg.UpdateUserTyping)
if !ok {
t.Fatalf("short update = %T, want *tg.UpdateUserTyping", short.Update)
}
if update.UserID != domain.ChatBotUserID {
t.Fatalf("typing user_id = %d, want ChatBot", update.UserID)
}
action, ok := update.Action.(*tg.SendMessageTextDraftAction)
if !ok {
t.Fatalf("typing action = %T, want *tg.SendMessageTextDraftAction", update.Action)
}
if action.RandomID != 4242 || action.Text.Text != "Hello from AI" {
t.Fatalf("draft action = random_id %d text %q", action.RandomID, action.Text.Text)
}
if snap := sessions.snapshot(); snap.userID != 1001 || snap.sessionID != 77 {
t.Fatalf("push target/excluded session = user %d session %d, want user 1001 session 77", snap.userID, snap.sessionID)
}
}

View file

@ -183,6 +183,11 @@ func tgWebPage(w domain.MessageWebPage) tg.WebPageClass {
}
}
}
if w.ComposeToneEmojiID != 0 {
page.SetAttributes([]tg.WebPageAttributeClass{
&tg.WebPageAttributeAiComposeTone{EmojiID: w.ComposeToneEmojiID},
})
}
return page
case domain.MessageWebPageStateEmpty:
page := &tg.WebPageEmpty{ID: w.ID}

View file

@ -98,6 +98,39 @@ func TestTgMessageMediaWebPageDone(t *testing.T) {
}
}
func TestTgMessageMediaWebPageAiComposeToneAttribute(t *testing.T) {
src := &domain.MessageMedia{
Kind: domain.MessageMediaKindWebPage,
WebPage: &domain.MessageWebPage{
State: domain.MessageWebPageStateDone,
ID: 123,
URL: "https://t.me/addstyle/ai-test",
DisplayURL: "t.me/addstyle/ai-test",
Hash: 7,
Type: "telegram_aicomposetone",
Title: "Sharp",
ComposeToneEmojiID: 99,
},
}
got := tgMessageMedia(jsonRoundTripMedia(t, src))
wrap, ok := got.(*tg.MessageMediaWebPage)
if !ok {
t.Fatalf("tgMessageMedia = %T, want *tg.MessageMediaWebPage", got)
}
page, ok := wrap.Webpage.(*tg.WebPage)
if !ok {
t.Fatalf("Webpage = %T, want *tg.WebPage", wrap.Webpage)
}
attrs, ok := page.GetAttributes()
if !ok || len(attrs) != 1 {
t.Fatalf("attributes = %#v ok=%v, want one", attrs, ok)
}
attr, ok := attrs[0].(*tg.WebPageAttributeAiComposeTone)
if !ok || attr.EmojiID != 99 {
t.Fatalf("attribute = %#v, want AiComposeTone emoji 99", attrs[0])
}
}
// TestTgMessageMediaWebPagePending 验证 pending 形态投影为 webPagePending{id,url,date}。
func TestTgMessageMediaWebPagePending(t *testing.T) {
src := &domain.MessageMedia{

View file

@ -67,6 +67,9 @@ func tgMessage(m domain.Message) tg.MessageClass {
if m.EditDate != 0 {
msg.SetEditDate(m.EditDate)
}
if m.HideEdited {
msg.SetEditHide(true)
}
if m.Silent {
msg.SetSilent(true)
}
@ -472,6 +475,12 @@ func tgMessageEntities(entities []domain.MessageEntity) []tg.MessageEntityClass
out = append(out, &tg.MessageEntityPhone{Offset: entity.Offset, Length: entity.Length})
case domain.MessageEntityBankCard:
out = append(out, &tg.MessageEntityBankCard{Offset: entity.Offset, Length: entity.Length})
case domain.MessageEntityDiffInsert:
out = append(out, &tg.MessageEntityDiffInsert{Offset: entity.Offset, Length: entity.Length})
case domain.MessageEntityDiffReplace:
out = append(out, &tg.MessageEntityDiffReplace{Offset: entity.Offset, Length: entity.Length, OldText: entity.OldText})
case domain.MessageEntityDiffDelete:
out = append(out, &tg.MessageEntityDiffDelete{Offset: entity.Offset, Length: entity.Length})
}
}
return out

View file

@ -658,12 +658,26 @@ type LangPackService interface {
GetStrings(ctx context.Context, langPack, langCode string, keys []string) (domain.LangPack, error)
}
// AIComposeService 抽象客户端输入框 AI 改写/润色与 aicompose tones 目录。
// 这里只使用 domain DTO;rpc 层负责 tg.TextWithEntities/InputAiComposeTone ↔ domain 转换。
type AIComposeService interface {
ListTones(ctx context.Context, userID, hash int64) (domain.AIComposeTones, bool, error)
GetTone(ctx context.Context, userID int64, ref domain.AIComposeToneRef) (domain.AIComposeTones, error)
CreateTone(ctx context.Context, input domain.AIComposeToneInput) (domain.AIComposeTone, error)
UpdateTone(ctx context.Context, update domain.AIComposeToneUpdate) (domain.AIComposeTone, error)
SaveTone(ctx context.Context, userID int64, ref domain.AIComposeToneRef, unsave bool) error
DeleteTone(ctx context.Context, userID int64, ref domain.AIComposeToneRef) error
GetToneExample(ctx context.Context, userID int64, ref domain.AIComposeToneRef, num int) (domain.AIComposeToneExample, error)
Compose(ctx context.Context, req domain.AIComposeRequest) (domain.AIComposeResult, error)
}
// Deps 按业务域注入服务接口。各域的 handler 注册见对应文件(auth.go / users.go / updates.go)。
type Deps struct {
Auth AuthService
Account AccountService
Privacy PrivacyService
Help HelpService
AICompose AIComposeService
Users UsersService
Updates UpdatesService
Contacts ContactsService

View file

@ -0,0 +1,34 @@
package rpc
import (
"testing"
"github.com/gotd/td/tg"
"telesrv/internal/domain"
)
func TestMessageProjectionSetsEditHide(t *testing.T) {
projected, ok := tgMessage(domain.Message{
ID: 5,
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: 1002},
Body: "streamed",
EditDate: 1700000400,
HideEdited: true,
}).(*tg.Message)
if !ok {
t.Fatalf("tgMessage = %T, want *tg.Message", projected)
}
if projected.EditDate != 1700000400 || !projected.EditHide {
t.Fatalf("projected edit fields = edit_date %d edit_hide %v", projected.EditDate, projected.EditHide)
}
plain := tgMessage(domain.Message{
ID: 6,
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: 1002},
Body: "edited",
EditDate: 1700000401,
}).(*tg.Message)
if plain.EditHide {
t.Fatal("plain edited message should not hide edited badge")
}
}

View file

@ -383,6 +383,9 @@ func (r *Router) webPagePreviewMedia(ctx context.Context, message string, entiti
// 的 20s,避免慢/挂上游把 RPC worker 钉死。命中(含负缓存的 empty)返回 ok=true,调用方据 state
// 决定;抓取失败返回 false。未启用返回 false。
func (r *Router) resolveWebPageForRequest(ctx context.Context, url string) (domain.MessageWebPage, bool) {
if page, ok := r.resolveAIComposeStyleWebPage(ctx, url); ok {
return page, true
}
if r.deps.Files == nil {
return domain.MessageWebPage{}, false
}

View file

@ -501,6 +501,7 @@ func (r *Router) registerMessages(d *tg.ServerDispatcher) {
d.OnMessagesGetSearchResultsCalendar(r.onMessagesGetSearchResultsCalendar)
d.OnMessagesGetSearchResultsPositions(r.onMessagesGetSearchResultsPositions)
d.OnMessagesSendReaction(r.onMessagesSendReaction)
d.OnMessagesComposeMessageWithAI(r.onMessagesComposeMessageWithAI)
// 语音转文字无识别后端:注册为显式失败(TRANSCRIPTION_FAILED),premium
// 客户端点击转录按钮得到优雅失败提示,而不是 NOT_IMPLEMENTED trace。
d.OnMessagesTranscribeAudio(func(ctx context.Context, req *tg.MessagesTranscribeAudioRequest) (*tg.MessagesTranscribedAudio, error) {

View file

@ -5,9 +5,13 @@ import (
"errors"
"testing"
"github.com/gotd/td/clock"
"github.com/gotd/td/tg"
"go.uber.org/zap/zaptest"
aiapp "telesrv/internal/app/ai"
"telesrv/internal/domain"
"telesrv/internal/store/memory"
)
// TestWebPagePreviewMedia 验证 getWebPagePreview 的 media 决策:done 卡片→messageMediaWebPage,
@ -71,6 +75,54 @@ func TestWebPagePreviewMedia(t *testing.T) {
})
}
func TestAIComposeStyleWebPagePreview(t *testing.T) {
const userID int64 = 1001
ctx := WithUserID(context.Background(), userID)
aiSvc := aiapp.NewService(memory.NewAIComposeStore())
tone, err := aiSvc.CreateTone(ctx, domain.AIComposeToneInput{
UserID: userID,
EmojiID: 12345,
Title: "Sharp",
Prompt: "Make it crisp.",
})
if err != nil {
t.Fatalf("CreateTone: %v", err)
}
r := New(Config{}, Deps{AICompose: aiSvc}, zaptest.NewLogger(t), clock.System)
link := "https://t.me/addstyle/" + tone.Slug
media := r.webPagePreviewMedia(ctx, "try "+link, nil)
wrap, ok := media.(*tg.MessageMediaWebPage)
if !ok {
t.Fatalf("media = %T, want *tg.MessageMediaWebPage", media)
}
page, ok := wrap.Webpage.(*tg.WebPage)
if !ok {
t.Fatalf("webpage = %T, want *tg.WebPage", wrap.Webpage)
}
if typ, ok := page.GetType(); !ok || typ != "telegram_aicomposetone" {
t.Fatalf("type = %q ok=%v, want telegram_aicomposetone", typ, ok)
}
if title, ok := page.GetTitle(); !ok || title != "Sharp" {
t.Fatalf("title = %q ok=%v, want Sharp", title, ok)
}
attrs, ok := page.GetAttributes()
if !ok || len(attrs) != 1 {
t.Fatalf("attributes = %#v ok=%v, want one", attrs, ok)
}
attr, ok := attrs[0].(*tg.WebPageAttributeAiComposeTone)
if !ok || attr.EmojiID != 12345 {
t.Fatalf("attribute = %#v, want ai compose tone emoji", attrs[0])
}
got := r.webPageForURL(ctx, "https://t.me/addstyle?slug="+tone.Slug, 0)
if page, ok := got.Webpage.(*tg.WebPage); !ok {
t.Fatalf("getWebPage webpage = %T, want *tg.WebPage", got.Webpage)
} else if typ, ok := page.GetType(); !ok || typ != "telegram_aicomposetone" {
t.Fatalf("getWebPage type = %q ok=%v", typ, ok)
}
}
func isEmptyMedia(m tg.MessageMediaClass) bool {
_, ok := m.(*tg.MessageMediaEmpty)
return ok

View file

@ -194,6 +194,7 @@ func TestStickersBotCreatePackLinkInstallIsolationSmoke(t *testing.T) {
Sessions: &captureSessions{},
}, zaptest.NewLogger(t), clock.System)
botsService.SetRouterHooks(r)
botsService.SetTextDraftPusher(r)
sendStickersBotText(t, r, alice, "/newpack", 9101)
waitForStickersReply(t, messageStore, alice.ID, "sticker pack")

View file

@ -111,6 +111,11 @@ func sliceUTF16(units []uint16, offset, length int) string {
//
// 未启用预览或 URL 不可规范化返回 nil(发送降级为无预览,不报错)。
func (r *Router) webPagePendingOrCachedMedia(ctx context.Context, rawURL string, invertMedia, forceLarge, forceSmall bool) *domain.MessageMedia {
if page, ok := r.resolveAIComposeStyleWebPage(ctx, rawURL); ok {
page.ForceLargeMedia = forceLarge
page.ForceSmallMedia = forceSmall
return &domain.MessageMedia{Kind: domain.MessageMediaKindWebPage, InvertMedia: invertMedia, WebPage: &page}
}
if r.deps.Files == nil || !r.deps.Files.WebPagePreviewEnabled() {
return nil
}