package laniakea import ( "errors" "strings" "testing" "git.scuroneko.dev/scuroneko/laniakea/tgapi" "git.scuroneko.dev/scuroneko/sneklog/v2" ) type sequenceDraftIDGenerator struct { ids []uint64 pos int } func (g *sequenceDraftIDGenerator) Next() uint64 { id := g.ids[g.pos] g.pos++ return id } func TestDraftFlushRequiresChatID(t *testing.T) { draft := NewRandomDraftProvider(&tgapi.API{}).NewDraft(tgapi.ParseNone) draft.Message = "hello" if err := draft.Flush(); !errors.Is(err, ErrDraftChatIDZero) { t.Fatalf("expected ErrDraftChatIDZero, got %v", err) } } func TestDraftFlushEmptyRemovesDraft(t *testing.T) { provider := NewLinearDraftProvider(nil, 0) draft := provider.NewDraft(tgapi.ParseNone) if err := draft.Flush(); err != nil { t.Fatalf("Flush returned error: %v", err) } if _, ok := provider.GetDraft(draft.ID); ok { t.Fatal("empty flushed draft remained in provider") } } func TestDraftFlushAllRemovesClearedDrafts(t *testing.T) { provider := NewLinearDraftProvider(nil, 0) draft := provider.NewDraft(tgapi.ParseNone) draft.Message = "discarded" draft.Clear() if err := provider.FlushAll(); err != nil { t.Fatalf("FlushAll returned error: %v", err) } if _, ok := provider.GetDraft(draft.ID); ok { t.Fatal("cleared draft remained in provider") } } func TestMsgContextNewDraftWorksWithoutLimiter(t *testing.T) { ctx := &MessageContext{ API: &tgapi.API{}, Msg: &tgapi.Message{ Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, }, Logger: sneklog.NewLogger(), draftProvider: NewRandomDraftProvider(&tgapi.API{}), } draft := ctx.NewDraft() if draft == nil { t.Fatal("expected draft") return } if draft.chatID != 42 { t.Fatalf("unexpected chat id: %d", draft.chatID) } } func TestDraftProviderSkipsZeroAndCollidingIDs(t *testing.T) { provider := &DraftProvider{ api: &tgapi.API{}, drafts: make(map[uint64]*Draft), generator: &sequenceDraftIDGenerator{ids: []uint64{0, 7, 7, 8}}, } first := provider.NewDraft(tgapi.ParseNone) second := provider.NewDraft(tgapi.ParseNone) if first.ID != 7 || second.ID != 8 { t.Fatalf("unexpected draft IDs: first=%d second=%d", first.ID, second.ID) } if got := len(provider.drafts); got != 2 { t.Fatalf("collision overwrote a draft: got %d drafts", got) } } func TestDraftReturnsErrorWhenAPIIsNil(t *testing.T) { draft := NewLinearDraftProvider(nil, 0).NewDraft(tgapi.ParseNone).SetChat(42, 0) if err := draft.Push("hello"); !errors.Is(err, ErrAPIIsNil) { t.Fatalf("expected ErrAPIIsNil from Push, got %v", err) } draft.Message = "hello" if err := draft.Flush(); !errors.Is(err, ErrAPIIsNil) { t.Fatalf("expected ErrAPIIsNil from Flush, got %v", err) } } func TestDraftFlushRejectsLongMessage(t *testing.T) { draft := NewRandomDraftProvider(&tgapi.API{}).NewDraft(tgapi.ParseNone).SetChat(42, 0) draft.Message = strings.Repeat("a", maxMessageTextLen+1) if err := draft.Flush(); !errors.Is(err, ErrMessageTooLong) { t.Fatalf("expected ErrMessageTooLong, got %v", err) } } func TestDraftPushRejectsLongMessage(t *testing.T) { draft := NewRandomDraftProvider(&tgapi.API{}).NewDraft(tgapi.ParseNone).SetChat(42, 0) if err := draft.Push(strings.Repeat("a", maxMessageTextLen+1)); !errors.Is(err, ErrMessageTooLong) { t.Fatalf("expected ErrMessageTooLong, got %v", err) } } // TestDraftPushLeavesMessageUnchangedOnValidationFailure covers the validation // order fix: when the candidate Message (current + new text) overflows the // Telegram limit, the existing Message must remain intact so callers can // recover and retry with a shorter payload instead of finding the draft in // a half-mutated state. func TestDraftPushLeavesMessageUnchangedOnValidationFailure(t *testing.T) { draft := NewRandomDraftProvider(&tgapi.API{}).NewDraft(tgapi.ParseNone).SetChat(42, 0) draft.Message = "hello" if err := draft.Push(strings.Repeat("a", maxMessageTextLen+1)); !errors.Is(err, ErrMessageTooLong) { t.Fatalf("expected ErrMessageTooLong, got %v", err) } if draft.Message != "hello" { t.Fatalf("expected draft Message to stay %q, got %q", "hello", draft.Message) } }