package socket import ( "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 validStartPaper() *altv1.StartPaperTradingRequest { return &altv1.StartPaperTradingRequest{ AccountId: "paper-1", Spec: &altv1.BacktestRunSpec{StrategyId: "strategy-v1"}, StartingCash: &altv1.Price{Currency: altv1.Currency_CURRENCY_KRW, Amount: &altv1.Decimal{Value: "10000000"}}, } } func TestHandleStartPaperTradingForwards(t *testing.T) { fake := &fakeWorkerClient{paperStartRes: &altv1.StartPaperTradingResponse{ State: &altv1.PaperTradingState{AccountId: "paper-1"}, }} resp, err := handleStartPaperTrading(fake, validStartPaper()) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.paperStartReq == nil { t.Fatal("expected request to be forwarded to worker") } if resp.GetState().GetAccountId() != "paper-1" { t.Errorf("unexpected account id: %q", resp.GetState().GetAccountId()) } } func TestHandleStartPaperTradingValidation(t *testing.T) { fake := &fakeWorkerClient{} cases := []*altv1.StartPaperTradingRequest{ {Spec: &altv1.BacktestRunSpec{StrategyId: "x"}, StartingCash: &altv1.Price{}}, // missing account_id {AccountId: "paper-1", StartingCash: &altv1.Price{}}, // missing spec {AccountId: "paper-1", Spec: &altv1.BacktestRunSpec{StrategyId: "x"}}, // missing starting_cash } for i, req := range cases { resp, err := handleStartPaperTrading(fake, req) if err != nil { t.Fatalf("case %d: expected typed validation response, got error: %v", i, err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) } if fake.paperStartReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleStartPaperTradingNilWorker(t *testing.T) { resp, err := handleStartPaperTrading(nil, validStartPaper()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) } func TestHandleStartPaperTradingConnectFailure(t *testing.T) { fake := &fakeWorkerClient{connectErr: workerclient.ErrUnavailable} resp, err := handleStartPaperTrading(fake, validStartPaper()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) if fake.paperStartReq != nil { t.Error("request must not be forwarded on connect failure") } } func TestHandleStartPaperTradingMapsTimeout(t *testing.T) { fake := &fakeWorkerClient{err: workerclient.ErrTimeout} resp, err := handleStartPaperTrading(fake, validStartPaper()) if err != nil { t.Fatalf("expected typed timeout response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorTimeout) } func TestHandleStartPaperTradingNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleStartPaperTrading(fake, validStartPaper()) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) if fake.paperStartReq == nil { t.Error("expected worker to be called before nil response was detected") } } func TestHandleGetPaperTradingStateForwards(t *testing.T) { fake := &fakeWorkerClient{paperStateRes: &altv1.GetPaperTradingStateResponse{ State: &altv1.PaperTradingState{AccountId: "paper-1"}, }} resp, err := handleGetPaperTradingState(fake, &altv1.GetPaperTradingStateRequest{AccountId: "paper-1"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.paperStateReq.GetAccountId() != "paper-1" { t.Errorf("account id not forwarded unchanged, got %q", fake.paperStateReq.GetAccountId()) } if resp.GetState().GetAccountId() != "paper-1" { t.Errorf("unexpected account id: %q", resp.GetState().GetAccountId()) } } func TestHandleGetPaperTradingStateRequiresAccountID(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleGetPaperTradingState(fake, &altv1.GetPaperTradingStateRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) if fake.paperStateReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleGetPaperTradingStateNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleGetPaperTradingState(fake, &altv1.GetPaperTradingStateRequest{AccountId: "paper-1"}) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) if fake.paperStateReq == nil { t.Error("expected request to be forwarded before nil response was detected") } } func TestPaperHandlersRegisteredInSession(t *testing.T) { registered := make(map[string]bool) for _, h := range sessionHandlers(nil) { registered[h.requestType] = true } for _, req := range []string{ protoSocket.TypeNameOf(&altv1.StartPaperTradingRequest{}), protoSocket.TypeNameOf(&altv1.GetPaperTradingStateRequest{}), } { if !registered[req] { t.Errorf("missing API handler registration for %q", req) } } } func TestPaperSocketWorkerUnavailableReturnsTypedError(t *testing.T) { fake := &fakeWorkerClient{err: workerclient.ErrUnavailable} client := startBacktestAPITestClient(t, fake) resp, err := protoSocket.SendRequestTyped[*altv1.StartPaperTradingRequest, *altv1.StartPaperTradingResponse]( &client.Communicator, validStartPaper(), 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) if fake.paperStartReq == nil { t.Fatal("valid start request should be forwarded before worker error is mapped") } } func TestPaperSocketValidationReturnsTypedError(t *testing.T) { fake := &fakeWorkerClient{} client := startBacktestAPITestClient(t, fake) resp, err := protoSocket.SendRequestTyped[*altv1.GetPaperTradingStateRequest, *altv1.GetPaperTradingStateResponse]( &client.Communicator, &altv1.GetPaperTradingStateRequest{}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) if fake.paperStateReq != nil { t.Fatal("worker must not be called for socket-level validation failures") } }