diff --git a/cmd/service/main.go b/cmd/service/main.go index 3106b69..e759be5 100644 --- a/cmd/service/main.go +++ b/cmd/service/main.go @@ -32,6 +32,7 @@ func main() { newGreetingDetector, newFloodDetector, newShortVoiceDetector, + newVoiceReplyWindowStore, newOwnerWhitelist, newBusinessConnectionStore, @@ -130,6 +131,10 @@ func newShortVoiceDetector(cfg *config.Config) *service.ShortVoiceDetector { }) } +func newVoiceReplyWindowStore(r *redis.Client) repository.VoiceReplyWindowStore { + return redisstore.NewVoiceReplyWindowStore(r) +} + func newOwnerWhitelist(cfg *config.Config) repository.OwnerWhitelist { return memory.NewOwnerWhitelist(cfg.AllowedOwners) } @@ -165,8 +170,9 @@ func newLLMClient(c deepseek.Config) repository.LLMClient { func newHandleBusinessMessageConfig(cfg *config.Config) handle_business_message.Config { return handle_business_message.Config{ - SystemPrompt: cfg.Bot.SystemPrompt, - ShortVoicePrompt: cfg.Bot.ShortVoicePrompt, + SystemPrompt: cfg.Bot.SystemPrompt, + ShortVoicePrompt: cfg.Bot.ShortVoicePrompt, + ShortVoiceResponseWindow: cfg.ShortVoice.ResponseWindow, } } diff --git a/internal/domain/repository/mock/voice_reply_window_store.go b/internal/domain/repository/mock/voice_reply_window_store.go new file mode 100644 index 0000000..b027616 --- /dev/null +++ b/internal/domain/repository/mock/voice_reply_window_store.go @@ -0,0 +1,71 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: voice_reply_window_store.go +// +// Generated by this command: +// +// mockgen -source=voice_reply_window_store.go -destination=mock/voice_reply_window_store.go -package=mock +// + +// Package mock is a generated GoMock package. +package mock + +import ( + context "context" + reflect "reflect" + time "time" + + gomock "go.uber.org/mock/gomock" +) + +// MockVoiceReplyWindowStore is a mock of VoiceReplyWindowStore interface. +type MockVoiceReplyWindowStore struct { + ctrl *gomock.Controller + recorder *MockVoiceReplyWindowStoreMockRecorder + isgomock struct{} +} + +// MockVoiceReplyWindowStoreMockRecorder is the mock recorder for MockVoiceReplyWindowStore. +type MockVoiceReplyWindowStoreMockRecorder struct { + mock *MockVoiceReplyWindowStore +} + +// NewMockVoiceReplyWindowStore creates a new mock instance. +func NewMockVoiceReplyWindowStore(ctrl *gomock.Controller) *MockVoiceReplyWindowStore { + mock := &MockVoiceReplyWindowStore{ctrl: ctrl} + mock.recorder = &MockVoiceReplyWindowStoreMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockVoiceReplyWindowStore) EXPECT() *MockVoiceReplyWindowStoreMockRecorder { + return m.recorder +} + +// Release mocks base method. +func (m *MockVoiceReplyWindowStore) Release(ctx context.Context, connectionID string, guestID int64) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Release", ctx, connectionID, guestID) + ret0, _ := ret[0].(error) + return ret0 +} + +// Release indicates an expected call of Release. +func (mr *MockVoiceReplyWindowStoreMockRecorder) Release(ctx, connectionID, guestID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Release", reflect.TypeOf((*MockVoiceReplyWindowStore)(nil).Release), ctx, connectionID, guestID) +} + +// TryEnter mocks base method. +func (m *MockVoiceReplyWindowStore) TryEnter(ctx context.Context, connectionID string, guestID int64, ttl time.Duration) (bool, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "TryEnter", ctx, connectionID, guestID, ttl) + ret0, _ := ret[0].(bool) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// TryEnter indicates an expected call of TryEnter. +func (mr *MockVoiceReplyWindowStoreMockRecorder) TryEnter(ctx, connectionID, guestID, ttl any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "TryEnter", reflect.TypeOf((*MockVoiceReplyWindowStore)(nil).TryEnter), ctx, connectionID, guestID, ttl) +} diff --git a/internal/domain/repository/voice_reply_window_store.go b/internal/domain/repository/voice_reply_window_store.go new file mode 100644 index 0000000..2d22ddc --- /dev/null +++ b/internal/domain/repository/voice_reply_window_store.go @@ -0,0 +1,18 @@ +package repository + +import ( + "context" + "time" +) + +//go:generate go tool mockgen -source=$GOFILE -destination=mock/$GOFILE -package=mock + +type VoiceReplyWindowStore interface { + // TryEnter atomically opens a reply window if it is not already open. + // Returns true if the window was just opened (the caller may reply), + // or false if the window is already open (the caller must skip). + TryEnter(ctx context.Context, connectionID string, guestID int64, ttl time.Duration) (bool, error) + + // Release drops the reservation if reply failed (DEL). + Release(ctx context.Context, connectionID string, guestID int64) error +} diff --git a/internal/domain/service/short_voice_detector.go b/internal/domain/service/short_voice_detector.go index b1cad27..120da01 100644 --- a/internal/domain/service/short_voice_detector.go +++ b/internal/domain/service/short_voice_detector.go @@ -19,7 +19,9 @@ func NewShortVoiceDetector(cfg ShortVoiceDetectorConfig) *ShortVoiceDetector { } } -func (d *ShortVoiceDetector) Detect(msg model.IncomingMessage) model.TriggerDecision { +func (d *ShortVoiceDetector) Detect( + msg model.IncomingMessage, +) model.TriggerDecision { if msg.Kind != model.MessageKindVoice { return model.TriggerDecision{Kind: model.TriggerKindNone} } diff --git a/internal/domain/service/short_voice_detector_test.go b/internal/domain/service/short_voice_detector_test.go index 561fa96..2b2d954 100644 --- a/internal/domain/service/short_voice_detector_test.go +++ b/internal/domain/service/short_voice_detector_test.go @@ -12,6 +12,10 @@ import ( func TestShortVoiceDetector_Detect(t *testing.T) { const maxDuration = 10 * time.Second + detector := service.NewShortVoiceDetector(service.ShortVoiceDetectorConfig{ + MaxDuration: maxDuration, + }) + tests := []struct { name string msg model.IncomingMessage @@ -26,7 +30,7 @@ func TestShortVoiceDetector_Detect(t *testing.T) { wantKind: model.TriggerKindShortVoice, }, { - name: "voice ровно на пороге — триггер (граничный случай)", + name: "voice ровно на пороге — триггер short_voice", msg: model.IncomingMessage{ Kind: model.MessageKindVoice, VoiceDuration: maxDuration, @@ -34,7 +38,7 @@ func TestShortVoiceDetector_Detect(t *testing.T) { wantKind: model.TriggerKindShortVoice, }, { - name: "voice длиннее порога — пропускаем", + name: "voice длиннее порога — без триггера", msg: model.IncomingMessage{ Kind: model.MessageKindVoice, VoiceDuration: 30 * time.Second, @@ -42,7 +46,7 @@ func TestShortVoiceDetector_Detect(t *testing.T) { wantKind: model.TriggerKindNone, }, { - name: "текстовое сообщение — детектор не реагирует", + name: "текстовое сообщение — без триггера", msg: model.IncomingMessage{ Kind: model.MessageKindText, Text: "привет", @@ -50,16 +54,12 @@ func TestShortVoiceDetector_Detect(t *testing.T) { wantKind: model.TriggerKindNone, }, { - name: "пустой Kind — детектор не реагирует", + name: "пустой Kind — без триггера", msg: model.IncomingMessage{}, wantKind: model.TriggerKindNone, }, } - detector := service.NewShortVoiceDetector(service.ShortVoiceDetectorConfig{ - MaxDuration: maxDuration, - }) - for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got := detector.Detect(tt.msg) diff --git a/internal/gateways/redis/voice_reply_window_store.go b/internal/gateways/redis/voice_reply_window_store.go new file mode 100644 index 0000000..535ad8f --- /dev/null +++ b/internal/gateways/redis/voice_reply_window_store.go @@ -0,0 +1,54 @@ +package redis + +import ( + "context" + "fmt" + "noirbot/internal/domain/repository" + "strconv" + "time" + + "github.com/redis/go-redis/v9" +) + +var _ repository.VoiceReplyWindowStore = (*VoiceReplyWindowStore)(nil) + +const voiceReplyKeyPrefix = "vrw:" + +type VoiceReplyWindowStore struct { + client *redis.Client +} + +func NewVoiceReplyWindowStore(client *redis.Client) *VoiceReplyWindowStore { + return &VoiceReplyWindowStore{ + client: client, + } +} + +func (s *VoiceReplyWindowStore) TryEnter( + ctx context.Context, + connectionID string, + guestID int64, + ttl time.Duration, +) (bool, error) { + key := voiceReplyKey(connectionID, guestID) + + ok, err := s.client.SetNX(ctx, key, "1", ttl).Result() + if err != nil { + return false, fmt.Errorf("voice reply window setnx: %w", err) + } + + return ok, nil +} + +func voiceReplyKey(connectionID string, guestID int64) string { + return voiceReplyKeyPrefix + connectionID + ":" + strconv.FormatInt(guestID, 10) +} + +func (s *VoiceReplyWindowStore) Release(ctx context.Context, connectionID string, guestID int64) error { + key := voiceReplyKey(connectionID, guestID) + if err := s.client.Del(ctx, key).Err(); err != nil { + return fmt.Errorf("voice reply window release: %w", err) + } + + return nil +} diff --git a/internal/gateways/redis/voice_reply_window_store_test.go b/internal/gateways/redis/voice_reply_window_store_test.go new file mode 100644 index 0000000..3df73fa --- /dev/null +++ b/internal/gateways/redis/voice_reply_window_store_test.go @@ -0,0 +1,126 @@ +package redis_test + +import ( + "context" + "testing" + "time" + + redisstore "noirbot/internal/gateways/redis" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func newVoiceReplyStore(t *testing.T) (*redisstore.VoiceReplyWindowStore, *miniredis.Miniredis) { + t.Helper() + + mr := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + + t.Cleanup(func() { _ = client.Close() }) + + return redisstore.NewVoiceReplyWindowStore(client), mr +} + +func TestVoiceReplyWindowStore_TryEnter_OpensThenBlocks(t *testing.T) { + store, _ := newVoiceReplyStore(t) + ctx := context.Background() + + const ( + connID = "conn-1" + guestID = int64(42) + ttl = time.Minute + ) + + opened, err := store.TryEnter(ctx, connID, guestID, ttl) + require.NoError(t, err) + require.True(t, opened, "first call must open the window") + + blocked, err := store.TryEnter(ctx, connID, guestID, ttl) + require.NoError(t, err) + require.False(t, blocked, "second call must see the window already open") +} + +func TestVoiceReplyWindowStore_TryEnter_TTLIsSet(t *testing.T) { + store, mr := newVoiceReplyStore(t) + ctx := context.Background() + + const ttl = 30 * time.Second + + _, err := store.TryEnter(ctx, "conn-1", 42, ttl) + require.NoError(t, err) + + keys := mr.Keys() + require.Len(t, keys, 1) + // miniredis rounds TTL to nearest second — check ≈ ttl + require.InDelta(t, ttl.Seconds(), mr.TTL(keys[0]).Seconds(), 1.0) +} + +func TestVoiceReplyWindowStore_TryEnter_KeyExpires(t *testing.T) { + store, mr := newVoiceReplyStore(t) + ctx := context.Background() + + const ttl = 3 * time.Second + + opened, err := store.TryEnter(ctx, "conn-1", 42, ttl) + require.NoError(t, err) + require.True(t, opened) + + // fast-forward past TTL + mr.FastForward(ttl + time.Second) + + // window must reopen after expiry + openedAgain, err := store.TryEnter(ctx, "conn-1", 42, ttl) + require.NoError(t, err) + require.True(t, openedAgain, "after TTL expires, window must reopen") +} + +func TestVoiceReplyWindowStore_TryEnter_ReleaseAllowsReentry(t *testing.T) { + store, _ := newVoiceReplyStore(t) + ctx := context.Background() + + const ( + connID = "conn-1" + guestID = int64(42) + ttl = time.Minute + ) + + opened, err := store.TryEnter(ctx, connID, guestID, ttl) + require.NoError(t, err) + require.True(t, opened) + + require.NoError(t, store.Release(ctx, connID, guestID)) + + openedAgain, err := store.TryEnter(ctx, connID, guestID, ttl) + require.NoError(t, err) + require.True(t, openedAgain, "after Release window must reopen") +} + +func TestVoiceReplyWindowStore_TryEnter_IsolatesPairs(t *testing.T) { + store, mr := newVoiceReplyStore(t) + ctx := context.Background() + ttl := time.Minute + + // conn-1, guest 42 — open + opened, err := store.TryEnter(ctx, "conn-1", 42, ttl) + require.NoError(t, err) + require.True(t, opened) + + // conn-1, guest 99 — different guest, must open + opened, err = store.TryEnter(ctx, "conn-1", 99, ttl) + require.NoError(t, err) + require.True(t, opened, "different guest must have own window") + + // conn-2, guest 42 — different connection, must open + opened, err = store.TryEnter(ctx, "conn-2", 42, ttl) + require.NoError(t, err) + require.True(t, opened, "different connection must have own window") + + // conn-1, guest 42 — still blocked + blocked, err := store.TryEnter(ctx, "conn-1", 42, ttl) + require.NoError(t, err) + require.False(t, blocked) + + require.Len(t, mr.Keys(), 3) +} diff --git a/internal/usecase/handle_business_message/errors.go b/internal/usecase/handle_business_message/errors.go index e800337..5e99be6 100644 --- a/internal/usecase/handle_business_message/errors.go +++ b/internal/usecase/handle_business_message/errors.go @@ -6,6 +6,7 @@ var ( ErrResolveOwner = errors.New("resolve owner failed") ErrWhitelistCheck = errors.New("whitelist check failed") ErrFloodDetect = errors.New("flood detection failed") + ErrVoiceWindow = errors.New("voice reply window failed") ErrLLMGenerate = errors.New("llm generate failed") ErrSend = errors.New("send reply failed") ) diff --git a/internal/usecase/handle_business_message/usecase.go b/internal/usecase/handle_business_message/usecase.go index b269dce..6238616 100644 --- a/internal/usecase/handle_business_message/usecase.go +++ b/internal/usecase/handle_business_message/usecase.go @@ -7,11 +7,13 @@ import ( "noirbot/internal/domain/model" "noirbot/internal/domain/repository" "noirbot/internal/domain/service" + "time" ) type Config struct { - SystemPrompt string - ShortVoicePrompt string + SystemPrompt string + ShortVoicePrompt string + ShortVoiceResponseWindow time.Duration } type Usecase struct { @@ -22,6 +24,7 @@ type Usecase struct { greetingDetector *service.GreetingDetector floodDetector *service.FloodDetector shortVoiceDetector *service.ShortVoiceDetector + voiceWindow repository.VoiceReplyWindowStore llm repository.LLMClient sender repository.BusinessSender log *slog.Logger @@ -40,6 +43,7 @@ func New( greetingDetector *service.GreetingDetector, floodDetector *service.FloodDetector, shortVoiceDetector *service.ShortVoiceDetector, + voiceWindow repository.VoiceReplyWindowStore, llm repository.LLMClient, sender repository.BusinessSender, log *slog.Logger, @@ -52,6 +56,7 @@ func New( greetingDetector: greetingDetector, floodDetector: floodDetector, shortVoiceDetector: shortVoiceDetector, + voiceWindow: voiceWindow, llm: llm, sender: sender, log: log.With("usecase", "handle_business_message"), @@ -88,6 +93,23 @@ func (uc *Usecase) Execute(ctx context.Context, msg model.IncomingMessage) error return nil } + shortVoice := decision.Kind == model.TriggerKindShortVoice + if shortVoice { + acquired, tryErr := uc.voiceWindow.TryEnter( + ctx, + msg.BusinessConnectionID, + msg.GuestID, + uc.cfg.ShortVoiceResponseWindow, + ) + if tryErr != nil { + return fmt.Errorf("%w: %w", ErrVoiceWindow, tryErr) + } + + if !acquired { + return nil + } + } + uc.log.InfoContext(ctx, "trigger fired", slog.String("kind", string(decision.Kind)), slog.String("reason", decision.Reason), @@ -108,18 +130,32 @@ func (uc *Usecase) Execute(ctx context.Context, msg model.IncomingMessage) error reply, err := uc.llm.Generate(ctx, in.SystemPrompt, in.UserText) if err != nil { + if shortVoice { + uc.releaseVoiceWindow(ctx, msg) + } + return fmt.Errorf("%w: %w", ErrLLMGenerate, err) } replyTarget.Text = reply if sndErr := uc.sender.Send(ctx, replyTarget); sndErr != nil { + if shortVoice { + uc.releaseVoiceWindow(ctx, msg) + } + return fmt.Errorf("%w: %w", ErrSend, sndErr) } return nil } +func (uc *Usecase) releaseVoiceWindow(ctx context.Context, msg model.IncomingMessage) { + if relErr := uc.voiceWindow.Release(ctx, msg.BusinessConnectionID, msg.GuestID); relErr != nil { + uc.log.WarnContext(ctx, "voice window release failed", "err", relErr) + } +} + func (uc *Usecase) resolveOwner(ctx context.Context, connectionID string) (model.Owner, error) { conn, ok, err := uc.connStore.Get(ctx, connectionID) if err != nil { diff --git a/internal/usecase/handle_business_message/usecase_test.go b/internal/usecase/handle_business_message/usecase_test.go index 38bf1e6..1ee05ae 100644 --- a/internal/usecase/handle_business_message/usecase_test.go +++ b/internal/usecase/handle_business_message/usecase_test.go @@ -5,6 +5,7 @@ import ( "errors" "log/slog" "noirbot/internal/domain/model" + "noirbot/internal/domain/repository" "noirbot/internal/domain/repository/mock" "noirbot/internal/domain/service" "testing" @@ -18,6 +19,7 @@ var ( errDeepseekTimeoutStub = errors.New("deepseek timeout") errTelegramRateLimitStub = errors.New("telegram 429") errTelegramDraftStub = errors.New("telegram draft unavailable") + errRedisRefusedStub = errors.New("redis connection refused") ) var ( @@ -44,6 +46,8 @@ var ( testReply = "Ну какой привет, пиши сразу, что тебе надо!" systemPrompt = "Отвечай как нуарный детектив, повидавший некоторое дерьмо" shortVoicePrompt = "Тебе пришло голосовое — отреагируй нуарно" + + responseWindow = 60 * time.Second ) func expectShowThinking(ctx context.Context, sender *mock.MockBusinessSender, msg model.IncomingMessage) { @@ -53,6 +57,82 @@ func expectShowThinking(ctx context.Context, sender *mock.MockBusinessSender, ms }).Return(nil) } +// usecaseMocks bundles the five collaborator mocks every test case builds. +type usecaseMocks struct { + whitelist *mock.MockOwnerWhitelist + connStore *mock.MockBusinessConnectionStore + accountReader *mock.MockBusinessAccountReader + llm *mock.MockLLMClient + sender *mock.MockBusinessSender +} + +func newMocks(ctrl *gomock.Controller) usecaseMocks { + return usecaseMocks{ + whitelist: mock.NewMockOwnerWhitelist(ctrl), + connStore: mock.NewMockBusinessConnectionStore(ctrl), + accountReader: mock.NewMockBusinessAccountReader(ctrl), + llm: mock.NewMockLLMClient(ctrl), + sender: mock.NewMockBusinessSender(ctrl), + } +} + +// expectAllowedOwner sets up the common cache-hit + whitelisted-owner path. +func (m usecaseMocks) expectAllowedOwner(ctx context.Context) { + m.connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) + m.whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) +} + +// expectTextReply sets up the full allowed-owner → think → generate → send happy +// path for testMsg. sendErr is what Send returns (nil for success). +func (m usecaseMocks) expectTextReply(ctx context.Context, sendErr error) { + m.expectAllowedOwner(ctx) + expectShowThinking(ctx, m.sender, testMsg) + m.llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) + m.sender.EXPECT().Send(ctx, gomock.Any()).Return(sendErr) +} + +func (m usecaseMocks) usecase(t *testing.T, voiceCooldown repository.VoiceReplyWindowStore) *Usecase { + t.Helper() + + return newUsecase(t, m.whitelist, m.connStore, m.accountReader, m.llm, m.sender, voiceCooldown) +} + +// mockVoiceCooldown returns a mock VoiceReplyWindowStore that expects +// TryEnter to NOT be called. Use for non-voice test cases. +func mockVoiceCooldown(ctrl *gomock.Controller) repository.VoiceReplyWindowStore { + return mock.NewMockVoiceReplyWindowStore(ctrl) +} + +// mockVoiceCooldownAcquired returns a mock that expects TryEnter → (true, nil). +func mockVoiceCooldownAcquired(ctrl *gomock.Controller) repository.VoiceReplyWindowStore { + store := mock.NewMockVoiceReplyWindowStore(ctrl) + store.EXPECT(). + TryEnter(gomock.Any(), "conn-1", int64(999), responseWindow). + Return(true, nil) + + return store +} + +// mockVoiceCooldownBlocked returns a mock that expects TryEnter → (false, nil). +func mockVoiceCooldownBlocked(ctrl *gomock.Controller) repository.VoiceReplyWindowStore { + store := mock.NewMockVoiceReplyWindowStore(ctrl) + store.EXPECT(). + TryEnter(gomock.Any(), "conn-1", int64(999), responseWindow). + Return(false, nil) + + return store +} + +// mockVoiceCooldownError returns a mock that expects TryEnter → error. +func mockVoiceCooldownError(ctrl *gomock.Controller) repository.VoiceReplyWindowStore { + store := mock.NewMockVoiceReplyWindowStore(ctrl) + store.EXPECT(). + TryEnter(gomock.Any(), "conn-1", int64(999), responseWindow). + Return(false, errRedisRefusedStub) + + return store +} + func newUsecase( t *testing.T, whitelist *mock.MockOwnerWhitelist, @@ -60,6 +140,7 @@ func newUsecase( accountReader *mock.MockBusinessAccountReader, llm *mock.MockLLMClient, sender *mock.MockBusinessSender, + voiceCooldown repository.VoiceReplyWindowStore, ) *Usecase { t.Helper() @@ -84,8 +165,9 @@ func newUsecase( return New( Config{ - SystemPrompt: systemPrompt, - ShortVoicePrompt: shortVoicePrompt, + SystemPrompt: systemPrompt, + ShortVoicePrompt: shortVoicePrompt, + ShortVoiceResponseWindow: responseWindow, }, whitelist, connStore, @@ -93,16 +175,45 @@ func newUsecase( greeting, flood, shortVoice, + voiceCooldown, llm, sender, slog.Default(), ) } -func TestUsecase_Execute(t *testing.T) { +func runUsecaseTests(t *testing.T, tests []struct { + name string + setup func(ctrl *gomock.Controller) *Usecase + msg model.IncomingMessage + wantErr error +}, +) { + t.Helper() + + ctx := context.Background() + + for i := range tests { + tt := &tests[i] + t.Run(tt.name, func(t *testing.T) { + ctrl := gomock.NewController(t) + uc := tt.setup(ctrl) + + err := uc.Execute(ctx, tt.msg) + + if tt.wantErr != nil { + require.ErrorIs(t, err, tt.wantErr) + } else { + require.NoError(t, err) + } + }) + } +} + +func TestUsecase_TextMessages(t *testing.T) { ctx := context.Background() - tests := []struct { + runUsecaseTests(t, []struct { name string setup func(ctrl *gomock.Controller) *Usecase msg model.IncomingMessage @@ -111,16 +222,11 @@ func TestUsecase_Execute(t *testing.T) { { name: "owner not in whitelist — LLM и sender не вызываются", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(false, nil) + m := newMocks(ctrl) + m.connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) + m.whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(false, nil) - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: testMsg, wantErr: nil, @@ -128,23 +234,17 @@ func TestUsecase_Execute(t *testing.T) { { name: "greeting match — LLM вызван, ответ отправлен", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - expectShowThinking(ctx, sender, testMsg) - llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) - sender.EXPECT().Send(ctx, model.ReplyDraft{ + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) + expectShowThinking(ctx, m.sender, testMsg) + m.llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) + m.sender.EXPECT().Send(ctx, model.ReplyDraft{ BusinessConnectionID: testMsg.BusinessConnectionID, GuestID: testMsg.GuestID, Text: testReply, }).Return(nil) - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: testMsg, wantErr: nil, @@ -152,16 +252,10 @@ func TestUsecase_Execute(t *testing.T) { { name: "длинное сообщение без приветствия — бот молчит", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: model.IncomingMessage{ BusinessConnectionID: "conn-1", @@ -172,22 +266,28 @@ func TestUsecase_Execute(t *testing.T) { }, wantErr: nil, }, + }) +} + +func TestUsecase_ErrorPropagation(t *testing.T) { + ctx := context.Background() + + runUsecaseTests(t, []struct { + name string + setup func(ctrl *gomock.Controller) *Usecase + msg model.IncomingMessage + wantErr error + }{ { name: "LLM вернул ошибку — возвращаем ErrLLMGenerate", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - expectShowThinking(ctx, sender, testMsg) - llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text). + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) + expectShowThinking(ctx, m.sender, testMsg) + m.llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text). Return("", errDeepseekTimeoutStub) - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: testMsg, wantErr: ErrLLMGenerate, @@ -195,64 +295,67 @@ func TestUsecase_Execute(t *testing.T) { { name: "sender вернул ошибку — возвращаем ErrSend", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - expectShowThinking(ctx, sender, testMsg) - llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) - sender.EXPECT().Send(ctx, gomock.Any()).Return(errTelegramRateLimitStub) - - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + m := newMocks(ctrl) + m.expectTextReply(ctx, errTelegramRateLimitStub) + + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: testMsg, wantErr: ErrSend, }, { - name: "cache miss — идём в accountReader, кешируем", + name: "show thinking failed — LLM и Send всё равно вызываются", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(model.BusinessConnection{}, false, nil) - accountReader.EXPECT().GetConnection(ctx, testConn.ID).Return(testConn, nil) - connStore.EXPECT().Put(ctx, testConn).Return(nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - expectShowThinking(ctx, sender, testMsg) - llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) - sender.EXPECT().Send(ctx, gomock.Any()).Return(nil) - - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) + m.sender.EXPECT().ShowThinking(ctx, model.ReplyDraft{ + BusinessConnectionID: testMsg.BusinessConnectionID, + GuestID: testMsg.GuestID, + }).Return(errTelegramDraftStub) + m.llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) + m.sender.EXPECT().Send(ctx, gomock.Any()).Return(nil) + + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: testMsg, wantErr: nil, }, { - name: "show thinking failed — LLM и Send всё равно вызываются", + name: "voice window store error → ErrVoiceWindow", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - sender.EXPECT().ShowThinking(ctx, model.ReplyDraft{ - BusinessConnectionID: testMsg.BusinessConnectionID, - GuestID: testMsg.GuestID, - }).Return(errTelegramDraftStub) - llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) - sender.EXPECT().Send(ctx, gomock.Any()).Return(nil) + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + return m.usecase(t, mockVoiceCooldownError(ctrl)) + }, + msg: testVoiceMsg, + wantErr: ErrVoiceWindow, + }, + }) +} + +func TestUsecase_EdgeCases(t *testing.T) { + ctx := context.Background() + + runUsecaseTests(t, []struct { + name string + setup func(ctrl *gomock.Controller) *Usecase + msg model.IncomingMessage + wantErr error + }{ + { + name: "cache miss — идём в accountReader, кешируем", + setup: func(ctrl *gomock.Controller) *Usecase { + m := newMocks(ctrl) + m.connStore.EXPECT().Get(ctx, testConn.ID).Return(model.BusinessConnection{}, false, nil) + m.accountReader.EXPECT().GetConnection(ctx, testConn.ID).Return(testConn, nil) + m.connStore.EXPECT().Put(ctx, testConn).Return(nil) + m.whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) + expectShowThinking(ctx, m.sender, testMsg) + m.llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) + m.sender.EXPECT().Send(ctx, gomock.Any()).Return(nil) + + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: testMsg, wantErr: nil, @@ -260,60 +363,84 @@ func TestUsecase_Execute(t *testing.T) { { name: "пустой whitelist (permissive) — любой owner проходит", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - expectShowThinking(ctx, sender, testMsg) - llm.EXPECT().Generate(ctx, systemPrompt, testMsg.Text).Return(testReply, nil) - sender.EXPECT().Send(ctx, gomock.Any()).Return(nil) - - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + m := newMocks(ctrl) + m.expectTextReply(ctx, nil) + + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: testMsg, wantErr: nil, }, + }) +} + +func TestUsecase_VoiceMessages(t *testing.T) { + ctx := context.Background() + + runUsecaseTests(t, []struct { + name string + setup func(ctrl *gomock.Controller) *Usecase + msg model.IncomingMessage + wantErr error + }{ { - name: "short voice ≤ порога — LLM вызван с short voice prompt и пустым userText", + name: "short voice ≤ порога + окно свободно — LLM вызван с short voice prompt", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) - - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) - expectShowThinking(ctx, sender, testVoiceMsg) - llm.EXPECT().Generate(ctx, shortVoicePrompt, "").Return(testReply, nil) - sender.EXPECT().Send(ctx, model.ReplyDraft{ + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) + expectShowThinking(ctx, m.sender, testVoiceMsg) + m.llm.EXPECT().Generate(ctx, shortVoicePrompt, "").Return(testReply, nil) + m.sender.EXPECT().Send(ctx, model.ReplyDraft{ BusinessConnectionID: testVoiceMsg.BusinessConnectionID, GuestID: testVoiceMsg.GuestID, Text: testReply, }).Return(nil) - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + return m.usecase(t, mockVoiceCooldownAcquired(ctrl)) }, msg: testVoiceMsg, wantErr: nil, }, { - name: "long voice > порога — бот молчит, LLM не вызывается", + name: "short voice + LLM error — Release вызывается, окно снимается", setup: func(ctrl *gomock.Controller) *Usecase { - whitelist := mock.NewMockOwnerWhitelist(ctrl) - connStore := mock.NewMockBusinessConnectionStore(ctrl) - accountReader := mock.NewMockBusinessAccountReader(ctrl) - llm := mock.NewMockLLMClient(ctrl) - sender := mock.NewMockBusinessSender(ctrl) + m := newMocks(ctrl) + voiceWindow := mock.NewMockVoiceReplyWindowStore(ctrl) + + m.expectAllowedOwner(ctx) + voiceWindow.EXPECT(). + TryEnter(gomock.Any(), "conn-1", int64(999), responseWindow). + Return(true, nil) + expectShowThinking(ctx, m.sender, testVoiceMsg) + m.llm.EXPECT().Generate(ctx, shortVoicePrompt, ""). + Return("", errDeepseekTimeoutStub) + voiceWindow.EXPECT(). + Release(gomock.Any(), "conn-1", int64(999)). + Return(nil) - connStore.EXPECT().Get(ctx, testConn.ID).Return(testConn, true, nil) - whitelist.EXPECT().IsAllowed(ctx, testConn.Owner.UserID).Return(true, nil) + return m.usecase(t, voiceWindow) + }, + msg: testVoiceMsg, + wantErr: ErrLLMGenerate, + }, + { + name: "short voice ≤ порога + окно занято — бот молчит", + setup: func(ctrl *gomock.Controller) *Usecase { + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) - return newUsecase(t, whitelist, connStore, accountReader, llm, sender) + return m.usecase(t, mockVoiceCooldownBlocked(ctrl)) + }, + msg: testVoiceMsg, + wantErr: nil, + }, + { + name: "long voice > порога — бот молчит, store и LLM не вызываются", + setup: func(ctrl *gomock.Controller) *Usecase { + m := newMocks(ctrl) + m.expectAllowedOwner(ctx) + // TryEnter is never called — gomock enforces this + return m.usecase(t, mockVoiceCooldown(ctrl)) }, msg: model.IncomingMessage{ BusinessConnectionID: "conn-1", @@ -324,20 +451,5 @@ func TestUsecase_Execute(t *testing.T) { }, wantErr: nil, }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ctrl := gomock.NewController(t) - uc := tt.setup(ctrl) - - err := uc.Execute(ctx, tt.msg) - - if tt.wantErr != nil { - require.ErrorIs(t, err, tt.wantErr) - } else { - require.NoError(t, err) - } - }) - } + }) } diff --git a/pkg/config/config.go b/pkg/config/config.go index e321fae..5f097d6 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -62,7 +62,8 @@ type FloodConfig struct { } type ShortVoiceConfig struct { - MaxDuration time.Duration `default:"10s" envconfig:"SHORT_VOICE_MAX_DURATION"` + MaxDuration time.Duration `default:"10s" envconfig:"SHORT_VOICE_MAX_DURATION"` + ResponseWindow time.Duration `default:"60s" envconfig:"SHORT_VOICE_RESPONSE_WINDOW"` } func Load() (*Config, error) { @@ -71,14 +72,18 @@ func Load() (*Config, error) { return nil, fmt.Errorf("load config: %w", err) } - if err := cfg.validate(); err != nil { + if err := cfg.validateFlood(); err != nil { + return nil, fmt.Errorf("validate config: %w", err) + } + + if err := cfg.validateShortVoice(); err != nil { return nil, fmt.Errorf("validate config: %w", err) } return cfg, nil } -func (c *Config) validate() error { +func (c *Config) validateFlood() error { if c.Flood.RedisTTL < c.Flood.WindowDuration { return fmt.Errorf("%w: ttl=%s, window=%s", ErrInvalidFloodTTL, @@ -89,3 +94,15 @@ func (c *Config) validate() error { return nil } + +func (c *Config) validateShortVoice() error { + if c.ShortVoice.ResponseWindow <= 0 || c.ShortVoice.MaxDuration <= 0 { + return fmt.Errorf("%w: response_window=%d, max_duration=%d", + ErrInvalidShortVoiceCfg, + c.ShortVoice.ResponseWindow, + c.ShortVoice.MaxDuration, + ) + } + + return nil +} diff --git a/pkg/config/errors.go b/pkg/config/errors.go index ba22870..6fd7b56 100644 --- a/pkg/config/errors.go +++ b/pkg/config/errors.go @@ -2,4 +2,7 @@ package config import "errors" -var ErrInvalidFloodTTL = errors.New("config: FLOOD_REDIS_TTL must be >= FLOOD_WINDOW") +var ( + ErrInvalidFloodTTL = errors.New("config: FLOOD_REDIS_TTL must be >= FLOOD_WINDOW") + ErrInvalidShortVoiceCfg = errors.New("config: SHORT_VOICE_MAX_DURATION and SHORT_VOICE_RESPONSE_WINDOW must be > 0") +)