package socket import ( "context" "testing" "time" altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" "git.toki-labs.com/toki/alt/services/api/internal/workerclient" protoSocket "git.toki-labs.com/toki/proto-socket/go" ) func validBarsRequest() *altv1.ListBarsRequest { return &altv1.ListBarsRequest{ InstrumentId: "KRX:005930", Timeframe: altv1.Timeframe_TIMEFRAME_DAILY, FromUnixMs: time.Date(2026, 5, 1, 0, 0, 0, 0, time.UTC).UnixMilli(), ToUnixMs: time.Date(2026, 5, 15, 0, 0, 0, 0, time.UTC).UnixMilli(), } } func validAggregateMonthlyRequest() *altv1.AggregateMonthlyBarsRequest { return &altv1.AggregateMonthlyBarsRequest{ Provider: "kis", SelectorKind: "watchlist", Market: altv1.Market_MARKET_KR, Venue: altv1.Venue_VENUE_KRX, Symbols: []string{"005930"}, FromYyyymmdd: "20260501", ToYyyymmdd: "20260531", } } func TestHandleListInstrumentsForwards(t *testing.T) { fake := &fakeWorkerClient{instRes: &altv1.ListInstrumentsResponse{ Instruments: []*altv1.Instrument{{Id: "KRX:005930", Symbol: "005930"}}, }} resp, err := handleListInstruments(fake, &altv1.ListInstrumentsRequest{ Market: altv1.Market_MARKET_KR, Provider: "kis", }) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.instReq == nil { t.Fatal("expected request to be forwarded to worker") } if fake.instReq.GetMarket() != altv1.Market_MARKET_KR || fake.instReq.GetProvider() != "kis" { t.Errorf("request was not forwarded unchanged: %+v", fake.instReq) } if len(resp.GetInstruments()) != 1 { t.Errorf("expected 1 instrument, got %d", len(resp.GetInstruments())) } } func TestHandleListInstrumentsForwardsUS(t *testing.T) { fake := &fakeWorkerClient{instRes: &altv1.ListInstrumentsResponse{ Instruments: []*altv1.Instrument{{Id: "NASDAQ:AAPL", Symbol: "AAPL"}}, }} resp, err := handleListInstruments(fake, &altv1.ListInstrumentsRequest{ Market: altv1.Market_MARKET_US, Provider: "kis", }) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.instReq == nil { t.Fatal("expected request to be forwarded to worker") } if fake.instReq.GetMarket() != altv1.Market_MARKET_US || fake.instReq.GetProvider() != "kis" { t.Errorf("request was not forwarded unchanged: %+v", fake.instReq) } if len(resp.GetInstruments()) != 1 { t.Errorf("expected 1 instrument, got %d", len(resp.GetInstruments())) } } func TestHandleListInstrumentsRejectsUnknownMarket(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleListInstruments(fake, &altv1.ListInstrumentsRequest{Market: altv1.Market(99)}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInvalidRequest) if fake.instReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleListInstrumentsMapsWorkerErrors(t *testing.T) { tests := []struct { name string err error code string }{ {"unavailable", workerclient.ErrUnavailable, marketErrorUnavailable}, {"timeout", workerclient.ErrTimeout, marketErrorTimeout}, {"unexpected", context.Canceled, marketErrorInternal}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { fake := &fakeWorkerClient{err: tt.err} resp, err := handleListInstruments(fake, &altv1.ListInstrumentsRequest{}) if err != nil { t.Fatalf("expected typed worker error response, got error: %v", err) } requireMarketError(t, resp.GetError(), tt.code) if fake.instReq == nil { t.Error("valid request should be forwarded before worker error is mapped") } }) } } func TestHandleListInstrumentsNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleListInstruments(fake, &altv1.ListInstrumentsRequest{}) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInternal) if fake.instReq == nil { t.Error("expected request to be forwarded before nil response was detected") } } func TestHandleListBarsForwards(t *testing.T) { fake := &fakeWorkerClient{barsRes: &altv1.ListBarsResponse{ Bars: []*altv1.Bar{{InstrumentId: "KRX:005930", Timeframe: altv1.Timeframe_TIMEFRAME_DAILY}}, }} resp, err := handleListBars(fake, validBarsRequest()) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.barsReq == nil { t.Fatal("expected request to be forwarded to worker") } if fake.barsReq.GetInstrumentId() != "KRX:005930" { t.Errorf("instrument id was not forwarded unchanged: %q", fake.barsReq.GetInstrumentId()) } if len(resp.GetBars()) != 1 { t.Errorf("expected 1 bar, got %d", len(resp.GetBars())) } } func TestHandleListBarsValidation(t *testing.T) { valid := validBarsRequest() tests := []struct { name string req *altv1.ListBarsRequest }{ {"missing instrument", &altv1.ListBarsRequest{Timeframe: valid.GetTimeframe(), FromUnixMs: valid.GetFromUnixMs(), ToUnixMs: valid.GetToUnixMs()}}, {"unknown timeframe", &altv1.ListBarsRequest{InstrumentId: valid.GetInstrumentId(), Timeframe: altv1.Timeframe(99), FromUnixMs: valid.GetFromUnixMs(), ToUnixMs: valid.GetToUnixMs()}}, {"missing from", &altv1.ListBarsRequest{InstrumentId: valid.GetInstrumentId(), Timeframe: valid.GetTimeframe(), ToUnixMs: valid.GetToUnixMs()}}, {"missing to", &altv1.ListBarsRequest{InstrumentId: valid.GetInstrumentId(), Timeframe: valid.GetTimeframe(), FromUnixMs: valid.GetFromUnixMs()}}, {"inverted range", &altv1.ListBarsRequest{InstrumentId: valid.GetInstrumentId(), Timeframe: valid.GetTimeframe(), FromUnixMs: valid.GetToUnixMs(), ToUnixMs: valid.GetFromUnixMs()}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleListBars(fake, tt.req) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInvalidRequest) if fake.barsReq != nil { t.Error("worker must not be called when validation fails") } }) } } func TestIsValidTimeframeAcceptsMonthlyAndMinuteBaseline(t *testing.T) { for _, tf := range []altv1.Timeframe{ altv1.Timeframe_TIMEFRAME_MONTHLY, altv1.Timeframe_TIMEFRAME_DAILY, altv1.Timeframe_TIMEFRAME_MINUTE_1, altv1.Timeframe_TIMEFRAME_MINUTE_5, } { if !isValidTimeframe(tf) { t.Errorf("expected timeframe %v to be valid", tf) } } } func TestHandleListBarsNilWorker(t *testing.T) { resp, err := handleListBars(nil, validBarsRequest()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorUnavailable) } func TestMarketSocketValidationReturnsTypedErrors(t *testing.T) { fake := &fakeWorkerClient{} client := startBacktestAPITestClient(t, fake) resp, err := protoSocket.SendRequestTyped[*altv1.ListBarsRequest, *altv1.ListBarsResponse]( &client.Communicator, &altv1.ListBarsRequest{}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInvalidRequest) if fake.barsReq != nil { t.Fatal("worker must not be called for socket-level validation failures") } } func TestMarketSocketListInstrumentsForwards(t *testing.T) { fake := &fakeWorkerClient{instRes: &altv1.ListInstrumentsResponse{ Instruments: []*altv1.Instrument{{Id: "KRX:005930", Symbol: "005930"}}, }} client := startBacktestAPITestClient(t, fake) resp, err := protoSocket.SendRequestTyped[*altv1.ListInstrumentsRequest, *altv1.ListInstrumentsResponse]( &client.Communicator, &altv1.ListInstrumentsRequest{Market: altv1.Market_MARKET_KR, Provider: "kis"}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return market response without timeout: %v", err) } if resp.GetError() != nil { t.Fatalf("unexpected market error: %+v", resp.GetError()) } if len(resp.GetInstruments()) != 1 { t.Fatalf("expected one instrument, got %d", len(resp.GetInstruments())) } if fake.instReq == nil || fake.instReq.GetProvider() != "kis" { t.Fatalf("expected request to be forwarded to worker, got %+v", fake.instReq) } } func TestMarketHandlersRegisteredInSession(t *testing.T) { registered := make(map[string]bool) for _, h := range sessionHandlers(nil) { registered[h.requestType] = true } required := []string{ protoSocket.TypeNameOf(&altv1.ListInstrumentsRequest{}), protoSocket.TypeNameOf(&altv1.ListBarsRequest{}), protoSocket.TypeNameOf(&altv1.ImportDailyBarsRequest{}), protoSocket.TypeNameOf(&altv1.AggregateMonthlyBarsRequest{}), } for _, req := range required { if !registered[req] { t.Errorf("missing API handler registration for %q", req) } } } func validImportRequest() *altv1.ImportDailyBarsRequest { return &altv1.ImportDailyBarsRequest{ Provider: "kis", SelectorKind: "watchlist", Market: altv1.Market_MARKET_KR, Venue: altv1.Venue_VENUE_KRX, Symbols: []string{"005930"}, FromYyyymmdd: "20260501", ToYyyymmdd: "20260515", } } func TestHandleImportDailyBarsForwards(t *testing.T) { fake := &fakeWorkerClient{importRes: &altv1.ImportDailyBarsResponse{ Provider: "kis", InstrumentCount: 1, BarCount: 10, }} resp, err := handleImportDailyBars(fake, validImportRequest()) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.importReq == nil { t.Fatal("expected request to be forwarded to worker") } if fake.importReq.GetProvider() != "kis" || len(fake.importReq.GetSymbols()) != 1 { t.Errorf("request was not forwarded unchanged: %+v", fake.importReq) } if resp.GetInstrumentCount() != 1 || resp.GetBarCount() != 10 { t.Errorf("unexpected response counts: %+v", resp) } } func TestHandleImportDailyBarsForwardsUS(t *testing.T) { fake := &fakeWorkerClient{importRes: &altv1.ImportDailyBarsResponse{ Provider: "kis", InstrumentCount: 2, BarCount: 20, }} req := &altv1.ImportDailyBarsRequest{ Provider: "kis", SelectorKind: "watchlist", Market: altv1.Market_MARKET_US, Venue: altv1.Venue_VENUE_NASDAQ, Symbols: []string{"AAPL", "MSFT"}, FromYyyymmdd: "20260501", ToYyyymmdd: "20260515", } resp, err := handleImportDailyBars(fake, req) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.importReq == nil { t.Fatal("expected request to be forwarded to worker") } if fake.importReq.GetMarket() != altv1.Market_MARKET_US || fake.importReq.GetVenue() != altv1.Venue_VENUE_NASDAQ { t.Errorf("request was not forwarded unchanged: %+v", fake.importReq) } if resp.GetInstrumentCount() != 2 || resp.GetBarCount() != 20 { t.Errorf("unexpected response counts: %+v", resp) } } func TestHandleImportDailyBarsValidation(t *testing.T) { // withImport starts from a fully valid request and mutates one field so each // case isolates a single validation failure. withImport := func(mut func(*altv1.ImportDailyBarsRequest)) *altv1.ImportDailyBarsRequest { req := validImportRequest() mut(req) return req } tests := []struct { name string req *altv1.ImportDailyBarsRequest }{ {"missing provider", withImport(func(r *altv1.ImportDailyBarsRequest) { r.Provider = "" })}, {"missing selector kind", withImport(func(r *altv1.ImportDailyBarsRequest) { r.SelectorKind = "" })}, {"missing symbols", withImport(func(r *altv1.ImportDailyBarsRequest) { r.Symbols = nil })}, {"unsupported market", withImport(func(r *altv1.ImportDailyBarsRequest) { r.Market = altv1.Market(99) })}, {"unsupported venue", withImport(func(r *altv1.ImportDailyBarsRequest) { r.Venue = altv1.Venue(99) })}, {"missing from date", withImport(func(r *altv1.ImportDailyBarsRequest) { r.FromYyyymmdd = "" })}, {"missing to date", withImport(func(r *altv1.ImportDailyBarsRequest) { r.ToYyyymmdd = "" })}, {"malformed from date", withImport(func(r *altv1.ImportDailyBarsRequest) { r.FromYyyymmdd = "2026-05-01" })}, {"malformed to date", withImport(func(r *altv1.ImportDailyBarsRequest) { r.ToYyyymmdd = "not-a-date" })}, {"reversed range", withImport(func(r *altv1.ImportDailyBarsRequest) { r.FromYyyymmdd, r.ToYyyymmdd = r.GetToYyyymmdd(), r.GetFromYyyymmdd() })}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleImportDailyBars(fake, tt.req) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInvalidRequest) if fake.importReq != nil { t.Error("worker must not be called when validation fails") } if fake.connectCount != 0 { t.Errorf("worker must not be connected when validation fails, got connectCount=%d", fake.connectCount) } }) } } func TestHandleImportDailyBarsNilWorker(t *testing.T) { resp, err := handleImportDailyBars(nil, validImportRequest()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorUnavailable) } func TestHandleImportDailyBarsNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleImportDailyBars(fake, validImportRequest()) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInternal) if fake.importReq == nil { t.Error("expected request to be forwarded before nil response was detected") } } func TestHandleImportDailyBarsMapsWorkerErrors(t *testing.T) { tests := []struct { name string err error code string }{ {"unavailable", workerclient.ErrUnavailable, marketErrorUnavailable}, {"timeout", workerclient.ErrTimeout, marketErrorTimeout}, {"unexpected", context.Canceled, marketErrorInternal}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { fake := &fakeWorkerClient{err: tt.err} resp, err := handleImportDailyBars(fake, validImportRequest()) if err != nil { t.Fatalf("expected typed worker error response, got error: %v", err) } requireMarketError(t, resp.GetError(), tt.code) if fake.importReq == nil { t.Error("valid request should be forwarded before worker error is mapped") } }) } } func TestHandleImportDailyBars_ConnectFailure(t *testing.T) { fake := &fakeWorkerClient{connectErr: workerclient.ErrUnavailable} resp, err := handleImportDailyBars(fake, validImportRequest()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorUnavailable) if fake.importReq != nil { t.Error("expected request not to be forwarded to worker on connect failure") } } func TestHandleAggregateMonthlyBarsForwards(t *testing.T) { fake := &fakeWorkerClient{monthlyRes: &altv1.AggregateMonthlyBarsResponse{ Provider: "kis", InstrumentCount: 1, SourceDailyBarCount: 21, MonthlyBarCount: 1, Provenance: []*altv1.MonthlyProvenance{ { InstrumentId: "KRX:005930", SourceDailyBarCount: 21, MonthlyBarCount: 1, SourceStartYyyymmdd: "20260501", SourceEndYyyymmdd: "20260531", AggregationRuleId: "monthly-deterministic-v1", SourceTimeframe: altv1.Timeframe_TIMEFRAME_DAILY, TargetTimeframe: altv1.Timeframe_TIMEFRAME_MONTHLY, }, }, }} req := validAggregateMonthlyRequest() resp, err := handleAggregateMonthlyBars(fake, req) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.monthlyReq != req { t.Fatalf("expected request to be forwarded unchanged, got %+v", fake.monthlyReq) } if resp.GetProvider() != "kis" || resp.GetInstrumentCount() != 1 || resp.GetSourceDailyBarCount() != 21 || resp.GetMonthlyBarCount() != 1 { t.Fatalf("unexpected aggregate response counts: %+v", resp) } if len(resp.GetProvenance()) != 1 { t.Fatalf("expected one provenance entry, got %d", len(resp.GetProvenance())) } } func TestHandleAggregateMonthlyBarsValidation(t *testing.T) { withMonthly := func(mut func(*altv1.AggregateMonthlyBarsRequest)) *altv1.AggregateMonthlyBarsRequest { req := validAggregateMonthlyRequest() mut(req) return req } tests := []struct { name string req *altv1.AggregateMonthlyBarsRequest }{ {"missing provider", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.Provider = "" })}, {"missing selector kind", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.SelectorKind = "" })}, {"missing symbols", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.Symbols = nil })}, {"unsupported market", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.Market = altv1.Market(99) })}, {"unsupported venue", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.Venue = altv1.Venue(99) })}, {"missing from date", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.FromYyyymmdd = "" })}, {"missing to date", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.ToYyyymmdd = "" })}, {"malformed from date", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.FromYyyymmdd = "2026-05-01" })}, {"malformed to date", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.ToYyyymmdd = "not-a-date" })}, {"reversed range", withMonthly(func(r *altv1.AggregateMonthlyBarsRequest) { r.FromYyyymmdd, r.ToYyyymmdd = r.GetToYyyymmdd(), r.GetFromYyyymmdd() })}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleAggregateMonthlyBars(fake, tt.req) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInvalidRequest) if fake.monthlyReq != nil { t.Error("worker must not be called when validation fails") } if fake.connectCount != 0 { t.Errorf("worker must not be connected when validation fails, got connectCount=%d", fake.connectCount) } }) } } func TestHandleAggregateMonthlyBarsMapsWorkerErrors(t *testing.T) { tests := []struct { name string err error code string }{ {"unavailable", workerclient.ErrUnavailable, marketErrorUnavailable}, {"timeout", workerclient.ErrTimeout, marketErrorTimeout}, {"unexpected", context.Canceled, marketErrorInternal}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { fake := &fakeWorkerClient{err: tt.err} resp, err := handleAggregateMonthlyBars(fake, validAggregateMonthlyRequest()) if err != nil { t.Fatalf("expected typed worker error response, got error: %v", err) } requireMarketError(t, resp.GetError(), tt.code) if fake.monthlyReq == nil { t.Error("valid request should be forwarded before worker error is mapped") } }) } } func TestHandleAggregateMonthlyBarsNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleAggregateMonthlyBars(fake, validAggregateMonthlyRequest()) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInternal) if fake.monthlyReq == nil { t.Error("expected request to be forwarded before nil response was detected") } } func TestHandleAggregateMonthlyBars_ConnectFailure(t *testing.T) { fake := &fakeWorkerClient{connectErr: workerclient.ErrUnavailable} resp, err := handleAggregateMonthlyBars(fake, validAggregateMonthlyRequest()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorUnavailable) if fake.monthlyReq != nil { t.Error("expected request not to be forwarded to worker on connect failure") } } func TestImportSocketValidationReturnsTypedErrors(t *testing.T) { fake := &fakeWorkerClient{} client := startBacktestAPITestClient(t, fake) resp, err := protoSocket.SendRequestTyped[*altv1.ImportDailyBarsRequest, *altv1.ImportDailyBarsResponse]( &client.Communicator, &altv1.ImportDailyBarsRequest{}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } requireMarketError(t, resp.GetError(), marketErrorInvalidRequest) if fake.importReq != nil { t.Fatal("worker must not be called for socket-level validation failures") } } func requireMarketError(t *testing.T, errInfo *altv1.ErrorInfo, code string) { t.Helper() if errInfo == nil { t.Fatalf("expected ErrorInfo code %q, got nil", code) } if errInfo.GetCode() != code { t.Fatalf("expected ErrorInfo code %q, got %q (%s)", code, errInfo.GetCode(), errInfo.GetMessage()) } if errInfo.GetMessage() == "" { t.Fatalf("expected ErrorInfo message for code %q", code) } } func TestHandleListInstruments_ConnectFailure(t *testing.T) { fake := &fakeWorkerClient{ connectErr: workerclient.ErrUnavailable, isConnected: false, } resp, err := handleListInstruments(fake, &altv1.ListInstrumentsRequest{ Market: altv1.Market_MARKET_KR, }) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireMarketError(t, resp.GetError(), marketErrorUnavailable) if fake.instReq != nil { t.Error("expected request not to be forwarded to worker on connect failure") } if fake.connectCount != 1 { t.Errorf("expected connectCount to be 1, got %d", fake.connectCount) } } func TestHandleListInstruments_ReconnectBehavior(t *testing.T) { fake := &fakeWorkerClient{ instRes: &altv1.ListInstrumentsResponse{ Instruments: []*altv1.Instrument{{Id: "KRX:005930"}}, }, isConnected: false, } resp, err := handleListInstruments(fake, &altv1.ListInstrumentsRequest{ Market: altv1.Market_MARKET_KR, }) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetError() != nil { t.Fatalf("unexpected handler error: %+v", resp.GetError()) } if fake.connectCount != 1 { t.Errorf("expected Connect to be called once, got %d", fake.connectCount) } if fake.instReq == nil { t.Error("expected request to be forwarded to worker") } if !fake.isConnected { t.Error("expected fake worker state to be connected") } }