From 3ee268816b79fde944255d0c317346b26a961369 Mon Sep 17 00:00:00 2001 From: toki Date: Sat, 30 May 2026 19:14:07 +0900 Subject: [PATCH] refactor: update API socket handlers and contract parsers - Add parser map for contract type resolution - Update socket handlers with new message routing - Add parser map tests - Remove outdated code review and plan docs for G07 --- .../CODE_REVIEW-cloud-G07.md | 25 ---- .../PLAN-cloud-G07.md | 132 ------------------ services/api/cmd/alt-api/main.go | 8 ++ services/api/internal/config/config.go | 2 + services/api/internal/config/config_test.go | 14 ++ services/api/internal/contracts/parser_map.go | 55 +++++--- .../api/internal/contracts/parser_map_test.go | 29 +++- services/api/internal/socket/handlers.go | 12 +- services/api/internal/socket/server.go | 5 + services/api/internal/socket/server_test.go | 22 +-- services/api/internal/workerclient/client.go | 124 ++++++++++++++++ services/worker/cmd/alt-worker/main.go | 23 +++ services/worker/internal/config/config.go | 18 +++ .../worker/internal/config/config_test.go | 81 ++++------- .../worker/internal/contracts/parser_map.go | 53 +++++++ services/worker/internal/socket/handlers.go | 56 ++++++++ services/worker/internal/socket/server.go | 39 ++++++ .../worker/internal/socket/server_test.go | 65 +++++++++ 18 files changed, 515 insertions(+), 248 deletions(-) delete mode 100644 agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/CODE_REVIEW-cloud-G07.md delete mode 100644 agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/PLAN-cloud-G07.md create mode 100644 services/api/internal/workerclient/client.go create mode 100644 services/worker/internal/contracts/parser_map.go create mode 100644 services/worker/internal/socket/handlers.go create mode 100644 services/worker/internal/socket/server.go create mode 100644 services/worker/internal/socket/server_test.go diff --git a/agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/CODE_REVIEW-cloud-G07.md b/agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/CODE_REVIEW-cloud-G07.md deleted file mode 100644 index 6c28c70..0000000 --- a/agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/CODE_REVIEW-cloud-G07.md +++ /dev/null @@ -1,25 +0,0 @@ - -# CODE_REVIEW-cloud-G07: Contracts and API Handler Registry Foundation - -## 구현 에이전트 소유 섹션 -- 구현 요약: 미작성 -- 변경 파일: 미작성 -- 실행한 검증/명령: 미작성 -- 남은 위험/후속 작업: 미작성 - -## 사용자 리뷰 요청 - -_기본값은 `없음`이다. 구현 중 사용자 결정, 외부 환경 준비, 또는 계획 범위 변경 없이는 안전하게 진행할 수 없으면 아래 항목을 실제 내용으로 교체하고, 구현을 중단한 뒤 active 파일을 그대로 둔 채 리뷰를 요청한다. code-review가 이 내용을 검증해 `USER_REVIEW.md`를 작성한다._ - -- 상태: 없음 -- 사유 유형: 없음 -- 결정 필요: 없음 -- 차단 근거: 없음 -- 실행한 검증/명령: 없음 -- 재개 조건: 없음 - -## 코드 리뷰어 소유 섹션 -- 리뷰 상태: 미작성 -- 주요 발견사항: 미작성 -- 테스트/검증 평가: 미작성 -- 판정: 미작성 diff --git a/agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/PLAN-cloud-G07.md b/agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/PLAN-cloud-G07.md deleted file mode 100644 index 320ab5d..0000000 --- a/agent-task/m-api-centered-proto-socket-rail/01_contracts_api_registry/PLAN-cloud-G07.md +++ /dev/null @@ -1,132 +0,0 @@ - -# PLAN-cloud-G07: Contracts and API Handler Registry Foundation - -## 이 파일을 읽는 구현 에이전트에게 -이 계획은 `API-Centered Proto-Socket Rail` 마일스톤 첫 번째 에픽의 큰 작업 중 `contracts`와 `api-hub`의 기반만 다룬다. 프로세스 간 연결을 바로 완성하려고 범위를 키우지 말고, API가 여러 요청 핸들러를 안정적으로 등록하고 계약 ID/파서가 빠지지 않게 만드는 기초 레일부터 깐다. - -## 배경 -현재 API proto-socket 서버는 `services/api/internal/socket/server.go:18`에서 만들어지고 `registerSessionHandlers`는 `HelloRequest`만 등록한다. 반면 계약에는 backtest/market 요청 메시지가 이미 있고, API parser map도 `services/api/internal/contracts/parser_map.go:11`에 여러 메시지를 등록하고 있다. 이후 작업이 API, worker, client로 나뉘므로 이 단계에서 핸들러 등록 구조와 계약 누락 확인 기준을 먼저 고정해야 한다. - -## 사용자 리뷰 요청 흐름 -구현 중 새 public contract를 추가해야 하는데 기존 요청/응답 의미를 바꿔야 한다면 중단하고 `CODE_REVIEW-cloud-G07.md`의 사용자 리뷰 요청 섹션을 채운다. 단순 additive proto 필드/메시지는 기존 호환성을 유지하는 범위에서 진행한다. - -## 분석 결과 -### 읽은 파일 -- `agent-roadmap/phase/operator-surface/milestones/api-centered-proto-socket-rail.md` -- `agent-ops/rules/project/rules.md` -- `agent-ops/rules/project/domain/api/rules.md` -- `services/api/internal/socket/server.go` -- `services/api/internal/socket/server_test.go` -- `services/api/internal/contracts/parser_map.go` -- `services/api/internal/contracts/parser_map_test.go` -- `packages/contracts/proto/alt/v1/common.proto` -- `packages/contracts/proto/alt/v1/market.proto` -- `packages/contracts/proto/alt/v1/backtest.proto` -- `packages/contracts/README.md` -- `../proto-socket/go/communicator.go` - -### 테스트 커버리지 공백 -- API socket server 테스트는 hello handler 중심이라 다중 핸들러 등록 실패, 중복 등록, parser 누락을 직접 잡지 못한다. -- 계약 parser map 테스트는 메시지 파싱을 확인하지만 API handler coverage와 연결되어 있지 않다. -- 로컬 테스트 실행은 `agent-test/local/rules.md`에 따라 금지되어 있으므로 원격 검증 환경에서만 실행한다. - -### 심볼 참조 -- `NewServer`: `services/api/internal/socket/server.go:18` -- `registerSessionHandlers`: `services/api/internal/socket/server.go:32` -- `NewParserMap`: `services/api/internal/contracts/parser_map.go:11` -- `StartBacktestRequest`: `packages/contracts/proto/alt/v1/backtest.proto:35` -- `ListBacktestRunsRequest`: `packages/contracts/proto/alt/v1/backtest.proto:95` -- `ListInstrumentsRequest`: `packages/contracts/proto/alt/v1/market.proto:39` -- `ListBarsRequest`: `packages/contracts/proto/alt/v1/market.proto:48` -- `AddRequestListenerTyped`: `../proto-socket/go/communicator.go` - -### 분할 판단 -이 작업은 API 내부 등록 구조와 계약 누락 점검까지만 맡는다. 실제 worker socket server/client 구현은 `02+01_worker_socket_rail`, backtest/market 비즈니스 연결은 `04+02_backtest_rail`, `05+02_market_rail`에서 처리한다. - -### 범위 결정 근거 -- 지금 바로 필요한 것은 API가 hello 외 요청을 받을 수 있는 구조적 자리다. -- worker process 연결 없이 API 핸들러 인터페이스와 parser/contract 점검을 먼저 끝내면 후속 구현 충돌이 줄어든다. -- 새 프로토콜 도입은 금지하고, 기존 proto-socket과 `packages/contracts/proto`만 사용한다. - -### 빌드 등급 -- Build lane: `cloud-G07` -- Review lane: `cloud-G07` -- 근거: API runtime protocol surface와 contracts registry를 건드리며, 로컬 검증이 금지되어 원격 테스트가 필요하다. - -## 구현 체크리스트 -### [API-1] API socket handler registry 분리 -문제: -`registerSessionHandlers`가 hello만 직접 등록하는 구조라 후속 요청을 추가할 때 socket server가 계속 비대해진다. - -해결 방법: -API socket 패키지 안에 handler registration 단위를 분리한다. 예시는 `type HandlerRegistrar interface`나 `RegisterHandlers(ctx, session)` 형태 중 기존 proto-socket 사용 방식에 맞는 최소 구조를 선택한다. hello handler도 새 등록 흐름을 통해 붙여 기존 동작을 유지한다. - -수정 파일 및 체크리스트: -- `services/api/internal/socket/server.go` -- 필요 시 `services/api/internal/socket/handlers.go` -- `services/api/internal/socket/server_test.go` -- [ ] hello handler가 새 registry 경유로 등록된다. -- [ ] 중복/누락 등록이 테스트에서 드러난다. -- [ ] public API 함수 시그니처 변경은 최소화한다. - -테스트 작성: -- hello request/response 기존 테스트 유지. -- registry에 복수 handler를 등록하는 테스트 추가. -- handler 등록 실패 또는 nil handler 방어 테스트를 추가할지 판단한다. - -중간 검증: -- 원격 검증 환경에서만 `go test ./services/api/...` 실행. - -### [API-2] Contracts/parser 누락 점검 고정 -문제: -API parser map은 market/backtest 메시지를 등록하지만, 실제 handler registry와 연결된 coverage가 약하다. - -해결 방법: -현재 milestone에서 API가 받을 요청 메시지 목록을 명시하고 parser map 테스트가 그 목록을 검증하게 한다. 기존 메시지로 부족한 경우에만 proto에 additive schema를 추가한다. - -수정 파일 및 체크리스트: -- `services/api/internal/contracts/parser_map.go` -- `services/api/internal/contracts/parser_map_test.go` -- 필요 시 `packages/contracts/proto/alt/v1/*.proto` -- 필요 시 generated contract files -- [ ] backtest start/list/detail/result/compare 요청/응답 parser가 모두 확인된다. -- [ ] market instruments/bars/status에 필요한 요청/응답 parser가 확인된다. -- [ ] schema 추가 시 Go/Dart generated artifacts 갱신 계획과 함께 처리한다. - -테스트 작성: -- parser map의 필수 message ID 목록 테스트. -- missing parser가 실패로 드러나는 table test. - -중간 검증: -- 원격 검증 환경에서만 `bin/contracts-check` 실행. -- 원격 검증 환경에서만 `go test ./services/api/...` 실행. - -### [API-3] 마일스톤 문서와 rule 간 용어 정렬 -문제: -API가 core/control plane 역할을 맡는다는 결정이 룰에는 반영됐지만 구현 계획의 용어도 같은 단어를 써야 후속 작업이 흔들리지 않는다. - -해결 방법: -구현 중 파일/패키지/테스트 이름에서 `core server`를 새 프로세스로 만들지 않는다. 필요한 naming은 `api hub`, `control plane`, `worker client`로 통일한다. - -수정 파일 및 체크리스트: -- 필요 시 `agent-roadmap/phase/operator-surface/milestones/api-centered-proto-socket-rail.md` -- 필요 시 domain rule 문서 -- [ ] 별도 core service 생성 없음. -- [ ] API 기준 runtime topology 유지. - -테스트 작성: -- 문서만 바꾸는 경우 테스트 없음. -- 코드 naming 변경 시 관련 API tests만 원격에서 실행. - -중간 검증: -- 원격 검증 환경에서만 관련 smoke 명령을 실행한다. - -## 수정 파일 요약 -- 예상 코드: `services/api/internal/socket/**`, `services/api/internal/contracts/**` -- 조건부 코드: `packages/contracts/proto/alt/v1/**`, generated contracts -- 예상 문서: 필요 시 milestone/rule 문구 정렬 - -## 최종 검증 -- 원격 검증 환경에서만 `go test ./services/api/...` -- 원격 검증 환경에서만 `bin/contracts-check` -모든 코드 변경 완료 후 반드시 `CODE_REVIEW-*-G??.md`의 구현 에이전트 소유 섹션을 채운다. 이 파일 작성이 구현의 마지막 단계다. diff --git a/services/api/cmd/alt-api/main.go b/services/api/cmd/alt-api/main.go index dfd2664..f27ce05 100644 --- a/services/api/cmd/alt-api/main.go +++ b/services/api/cmd/alt-api/main.go @@ -19,6 +19,10 @@ func main() { cfg := config.Load() server := socket.NewServer(cfg) + go func() { + _ = socket.Worker.Connect(ctx) + }() + if err := server.Start(ctx); err != nil { slog.Error("failed to start api socket server", "error", err) os.Exit(1) @@ -30,4 +34,8 @@ func main() { slog.Error("failed to stop api socket server", "error", err) os.Exit(1) } + + if err := socket.Worker.Close(); err != nil { + slog.Error("failed to close worker socket client", "error", err) + } } diff --git a/services/api/internal/config/config.go b/services/api/internal/config/config.go index bfa03e7..8b9339a 100644 --- a/services/api/internal/config/config.go +++ b/services/api/internal/config/config.go @@ -13,6 +13,7 @@ type Config struct { HeartbeatIntervalSec int HeartbeatWaitSec int WSOriginPatterns []string + WorkerSocketURL string } func Load() Config { @@ -23,6 +24,7 @@ func Load() Config { HeartbeatIntervalSec: getenvInt("ALT_SOCKET_HEARTBEAT_INTERVAL_SEC", 30), HeartbeatWaitSec: getenvInt("ALT_SOCKET_HEARTBEAT_WAIT_SEC", 10), WSOriginPatterns: getenvList("ALT_API_WS_ORIGIN_PATTERNS", "localhost:*,127.0.0.1:*"), + WorkerSocketURL: getenv("ALT_WORKER_SOCKET_URL", "ws://127.0.0.1:8081/socket"), } } diff --git a/services/api/internal/config/config_test.go b/services/api/internal/config/config_test.go index c25eec2..af12308 100644 --- a/services/api/internal/config/config_test.go +++ b/services/api/internal/config/config_test.go @@ -31,3 +31,17 @@ func TestLoadParsesWebSocketOriginPatterns(t *testing.T) { t.Fatalf("WSOriginPatterns = %v, want %v", cfg.WSOriginPatterns, want) } } + +func TestLoadWorkerSocketURL(t *testing.T) { + t.Setenv("ALT_WORKER_SOCKET_URL", "ws://127.0.0.1:9090/custom") + cfg := Load() + if cfg.WorkerSocketURL != "ws://127.0.0.1:9090/custom" { + t.Errorf("expected WorkerSocketURL to be ws://127.0.0.1:9090/custom, got %q", cfg.WorkerSocketURL) + } + + t.Setenv("ALT_WORKER_SOCKET_URL", "") + cfg = Load() + if cfg.WorkerSocketURL != "ws://127.0.0.1:8081/socket" { + t.Errorf("expected default WorkerSocketURL to be ws://127.0.0.1:8081/socket, got %q", cfg.WorkerSocketURL) + } +} diff --git a/services/api/internal/contracts/parser_map.go b/services/api/internal/contracts/parser_map.go index de1f6d2..229b4a4 100644 --- a/services/api/internal/contracts/parser_map.go +++ b/services/api/internal/contracts/parser_map.go @@ -7,29 +7,44 @@ import ( protoSocket "git.toki-labs.com/toki/proto-socket/go" ) +// messageFactories lists every ALT protobuf message the API parser map must +// decode. A single factory list keeps request/response pairs together so a new +// contract message cannot be registered through one code path and missed in +// another, which is the drift this milestone's contract check guards against. +func messageFactories() []func() proto.Message { + return []func() proto.Message{ + func() proto.Message { return &altv1.HelloRequest{} }, + func() proto.Message { return &altv1.HelloResponse{} }, + // market read surface: instruments and bars + func() proto.Message { return &altv1.ListInstrumentsRequest{} }, + func() proto.Message { return &altv1.ListInstrumentsResponse{} }, + func() proto.Message { return &altv1.ListBarsRequest{} }, + func() proto.Message { return &altv1.ListBarsResponse{} }, + // backtest surface: start / list / detail / result / compare + func() proto.Message { return &altv1.StartBacktestRequest{} }, + func() proto.Message { return &altv1.StartBacktestResponse{} }, + func() proto.Message { return &altv1.GetBacktestRunRequest{} }, + func() proto.Message { return &altv1.GetBacktestRunResponse{} }, + func() proto.Message { return &altv1.GetBacktestResultRequest{} }, + func() proto.Message { return &altv1.GetBacktestResultResponse{} }, + func() proto.Message { return &altv1.BacktestResult{} }, + func() proto.Message { return &altv1.ListBacktestRunsRequest{} }, + func() proto.Message { return &altv1.ListBacktestRunsResponse{} }, + func() proto.Message { return &altv1.GetBacktestRunDetailRequest{} }, + func() proto.Message { return &altv1.GetBacktestRunDetailResponse{} }, + func() proto.Message { return &altv1.CompareBacktestRunsRequest{} }, + func() proto.Message { return &altv1.CompareBacktestRunsResponse{} }, + } +} + // ParserMap returns a fresh ParserMap populated with all ALT API messages. func ParserMap() protoSocket.ParserMap { - return protoSocket.ParserMap{ - protoSocket.TypeNameOf(&altv1.HelloRequest{}): parserFor(func() proto.Message { return &altv1.HelloRequest{} }), - protoSocket.TypeNameOf(&altv1.HelloResponse{}): parserFor(func() proto.Message { return &altv1.HelloResponse{} }), - protoSocket.TypeNameOf(&altv1.ListInstrumentsRequest{}): parserFor(func() proto.Message { return &altv1.ListInstrumentsRequest{} }), - protoSocket.TypeNameOf(&altv1.ListInstrumentsResponse{}): parserFor(func() proto.Message { return &altv1.ListInstrumentsResponse{} }), - protoSocket.TypeNameOf(&altv1.ListBarsRequest{}): parserFor(func() proto.Message { return &altv1.ListBarsRequest{} }), - protoSocket.TypeNameOf(&altv1.ListBarsResponse{}): parserFor(func() proto.Message { return &altv1.ListBarsResponse{} }), - protoSocket.TypeNameOf(&altv1.StartBacktestRequest{}): parserFor(func() proto.Message { return &altv1.StartBacktestRequest{} }), - protoSocket.TypeNameOf(&altv1.StartBacktestResponse{}): parserFor(func() proto.Message { return &altv1.StartBacktestResponse{} }), - protoSocket.TypeNameOf(&altv1.GetBacktestRunRequest{}): parserFor(func() proto.Message { return &altv1.GetBacktestRunRequest{} }), - protoSocket.TypeNameOf(&altv1.GetBacktestRunResponse{}): parserFor(func() proto.Message { return &altv1.GetBacktestRunResponse{} }), - protoSocket.TypeNameOf(&altv1.GetBacktestResultRequest{}): parserFor(func() proto.Message { return &altv1.GetBacktestResultRequest{} }), - protoSocket.TypeNameOf(&altv1.GetBacktestResultResponse{}): parserFor(func() proto.Message { return &altv1.GetBacktestResultResponse{} }), - protoSocket.TypeNameOf(&altv1.BacktestResult{}): parserFor(func() proto.Message { return &altv1.BacktestResult{} }), - protoSocket.TypeNameOf(&altv1.ListBacktestRunsRequest{}): parserFor(func() proto.Message { return &altv1.ListBacktestRunsRequest{} }), - protoSocket.TypeNameOf(&altv1.ListBacktestRunsResponse{}): parserFor(func() proto.Message { return &altv1.ListBacktestRunsResponse{} }), - protoSocket.TypeNameOf(&altv1.GetBacktestRunDetailRequest{}): parserFor(func() proto.Message { return &altv1.GetBacktestRunDetailRequest{} }), - protoSocket.TypeNameOf(&altv1.GetBacktestRunDetailResponse{}): parserFor(func() proto.Message { return &altv1.GetBacktestRunDetailResponse{} }), - protoSocket.TypeNameOf(&altv1.CompareBacktestRunsRequest{}): parserFor(func() proto.Message { return &altv1.CompareBacktestRunsRequest{} }), - protoSocket.TypeNameOf(&altv1.CompareBacktestRunsResponse{}): parserFor(func() proto.Message { return &altv1.CompareBacktestRunsResponse{} }), + factories := messageFactories() + pm := make(protoSocket.ParserMap, len(factories)) + for _, factory := range factories { + pm[protoSocket.TypeNameOf(factory())] = parserFor(factory) } + return pm } func parserFor(factory func() proto.Message) func([]byte) (proto.Message, error) { diff --git a/services/api/internal/contracts/parser_map_test.go b/services/api/internal/contracts/parser_map_test.go index 9e9ab72..f7b099c 100644 --- a/services/api/internal/contracts/parser_map_test.go +++ b/services/api/internal/contracts/parser_map_test.go @@ -9,16 +9,20 @@ import ( altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" ) -func TestParserMapIncludesAltMessages(t *testing.T) { - pm := ParserMap() - - messages := []proto.Message{ +// requiredAPIMessages is the milestone contract surface the API parser map must +// cover. It is declared independently from messageFactories so that removing a +// message from the parser map is caught here as a missing parser instead of +// silently shrinking both sides together. +func requiredAPIMessages() []proto.Message { + return []proto.Message{ &altv1.HelloRequest{}, &altv1.HelloResponse{}, + // market instruments / bars &altv1.ListInstrumentsRequest{}, &altv1.ListInstrumentsResponse{}, &altv1.ListBarsRequest{}, &altv1.ListBarsResponse{}, + // backtest start / list / detail / result / compare &altv1.StartBacktestRequest{}, &altv1.StartBacktestResponse{}, &altv1.GetBacktestRunRequest{}, @@ -33,8 +37,12 @@ func TestParserMapIncludesAltMessages(t *testing.T) { &altv1.CompareBacktestRunsRequest{}, &altv1.CompareBacktestRunsResponse{}, } +} - for _, msg := range messages { +func TestParserMapIncludesAltMessages(t *testing.T) { + pm := ParserMap() + + for _, msg := range requiredAPIMessages() { typeName := protoSocket.TypeNameOf(msg) t.Run(typeName, func(t *testing.T) { parser := pm[typeName] @@ -58,3 +66,14 @@ func TestParserMapIncludesAltMessages(t *testing.T) { }) } } + +// TestParserMapReportsMissingParser confirms a missing parser is observable as +// a nil lookup, so the coverage assertions above fail loudly when a required +// message is dropped from the parser map. +func TestParserMapReportsMissingParser(t *testing.T) { + pm := ParserMap() + + if parser := pm["alt.v1.UnregisteredProbeMessage"]; parser != nil { + t.Fatal("expected nil parser for an unregistered message type") + } +} diff --git a/services/api/internal/socket/handlers.go b/services/api/internal/socket/handlers.go index cd3acc2..881babb 100644 --- a/services/api/internal/socket/handlers.go +++ b/services/api/internal/socket/handlers.go @@ -29,10 +29,16 @@ func sessionHandlers() []sessionHandler { } // registerSessionHandlers wires every session handler onto a newly connected -// client. nil registrars are skipped so a malformed registry entry cannot -// panic the connection setup path. +// client. It is the OnClientConnected entrypoint for the API socket server. func registerSessionHandlers(client *protoSocket.WsClient) { - for _, handler := range sessionHandlers() { + registerHandlers(client, sessionHandlers()) +} + +// registerHandlers attaches the given handlers onto a client. nil registrars +// are skipped so a malformed registry entry cannot panic the connection setup +// path. +func registerHandlers(client *protoSocket.WsClient, handlers []sessionHandler) { + for _, handler := range handlers { if handler.register == nil { continue } diff --git a/services/api/internal/socket/server.go b/services/api/internal/socket/server.go index d352886..86c7763 100644 --- a/services/api/internal/socket/server.go +++ b/services/api/internal/socket/server.go @@ -6,6 +6,7 @@ import ( "git.toki-labs.com/toki/alt/services/api/internal/config" apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts" + "git.toki-labs.com/toki/alt/services/api/internal/workerclient" ) const ( @@ -14,7 +15,11 @@ const ( defaultAltProtocolVersion = "alt.v1" ) +var Worker workerclient.WorkerClient + func NewServer(cfg config.Config) *protoSocket.WsServer { + Worker = workerclient.New(cfg.WorkerSocketURL) + options := protoSocket.WsServerOptions{} if len(cfg.WSOriginPatterns) > 0 { options.AcceptOptions = &websocket.AcceptOptions{ diff --git a/services/api/internal/socket/server_test.go b/services/api/internal/socket/server_test.go index 6e112d6..c71473a 100644 --- a/services/api/internal/socket/server_test.go +++ b/services/api/internal/socket/server_test.go @@ -95,20 +95,26 @@ func TestSessionHandlersCoverRequiredRequests(t *testing.T) { } } -func TestRegisterSessionHandlersSkipsNilRegistrar(t *testing.T) { - // registerSessionHandlers must tolerate a malformed registry entry instead - // of panicking during connection setup. +func TestRegisterHandlersSkipsNilRegistrar(t *testing.T) { + // registerHandlers must tolerate a malformed registry entry instead of + // panicking during connection setup. A nil registrar is skipped without + // dereferencing the client. defer func() { if r := recover(); r != nil { - t.Fatalf("registerSessionHandlers panicked on nil registrar: %v", r) + t.Fatalf("registerHandlers panicked on nil registrar: %v", r) } }() - handler := sessionHandler{requestType: "alt.v1.NilRegistrarProbe", register: nil} - if handler.register != nil { - t.Fatal("expected nil registrar for probe handler") + called := false + handlers := []sessionHandler{ + {requestType: "alt.v1.NilRegistrarProbe", register: nil}, + {requestType: "alt.v1.LiveRegistrarProbe", register: func(*protoSocket.WsClient) { called = true }}, + } + registerHandlers(nil, handlers) + + if !called { + t.Fatal("expected non-nil registrar to run after nil registrar was skipped") } - registerSessionHandlers(nil) } func freeTCPPort(t *testing.T) int { diff --git a/services/api/internal/workerclient/client.go b/services/api/internal/workerclient/client.go new file mode 100644 index 0000000..027a646 --- /dev/null +++ b/services/api/internal/workerclient/client.go @@ -0,0 +1,124 @@ +package workerclient + +import ( + "context" + "errors" + "fmt" + "net" + "net/url" + "strconv" + "sync" + "time" + + altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" + apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts" + protoSocket "git.toki-labs.com/toki/proto-socket/go" +) + +var ( + ErrUnavailable = errors.New("worker is not available") + ErrTimeout = errors.New("worker request timeout") +) + +type WorkerClient interface { + Connect(ctx context.Context) error + Close() error + Hello(ctx context.Context, req *altv1.HelloRequest) (*altv1.HelloResponse, error) +} + +type socketClient struct { + socketURL string + mu sync.RWMutex + wsClient *protoSocket.WsClient +} + +func New(socketURL string) WorkerClient { + return &socketClient{ + socketURL: socketURL, + } +} + +func (c *socketClient) Connect(ctx context.Context) error { + c.mu.Lock() + defer c.mu.Unlock() + + if c.wsClient != nil && c.wsClient.IsAlive() { + return nil + } + + u, err := url.Parse(c.socketURL) + if err != nil { + return fmt.Errorf("invalid worker socket URL %q: %w", c.socketURL, err) + } + + host, portStr, err := net.SplitHostPort(u.Host) + if err != nil { + host = u.Host + if u.Scheme == "wss" { + portStr = "443" + } else { + portStr = "80" + } + } + + port, err := strconv.Atoi(portStr) + if err != nil { + return fmt.Errorf("invalid worker port %q: %w", portStr, err) + } + + path := u.Path + if path == "" { + path = "/" + } + + var wsClient *protoSocket.WsClient + if u.Scheme == "wss" { + wsClient, err = protoSocket.DialWssWithHeartbeat(ctx, host, port, path, nil, 30, 10, apiContracts.ParserMap()) + } else { + wsClient, err = protoSocket.DialWsWithHeartbeat(ctx, host, port, path, 30, 10, apiContracts.ParserMap()) + } + + if err != nil { + return fmt.Errorf("%w: %v", ErrUnavailable, err) + } + + c.wsClient = wsClient + return nil +} + +func (c *socketClient) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + + if c.wsClient != nil { + err := c.wsClient.Close() + c.wsClient = nil + return err + } + return nil +} + +func (c *socketClient) Hello(ctx context.Context, req *altv1.HelloRequest) (*altv1.HelloResponse, error) { + c.mu.RLock() + client := c.wsClient + c.mu.RUnlock() + + if client == nil || !client.IsAlive() { + return nil, ErrUnavailable + } + + timeout := 5 * time.Second + if dl, ok := ctx.Deadline(); ok { + timeout = time.Until(dl) + } + + res, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, req, timeout) + if err != nil { + if errors.Is(err, protoSocket.ErrNotConnected) { + return nil, ErrUnavailable + } + return nil, fmt.Errorf("%w: %v", ErrTimeout, err) + } + + return res, nil +} diff --git a/services/worker/cmd/alt-worker/main.go b/services/worker/cmd/alt-worker/main.go index 48d2277..a125bc0 100644 --- a/services/worker/cmd/alt-worker/main.go +++ b/services/worker/cmd/alt-worker/main.go @@ -1,19 +1,42 @@ package main import ( + "context" + "fmt" "log/slog" + "os" + "os/signal" + "syscall" "git.toki-labs.com/toki/alt/services/worker/internal/config" "git.toki-labs.com/toki/alt/services/worker/internal/jobs" + "git.toki-labs.com/toki/alt/services/worker/internal/socket" ) func main() { + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + cfg := config.Load() runner := jobs.NewRunner() jobs.RegisterBuiltins(runner) + + server := socket.NewServer(cfg) + if err := server.Start(ctx); err != nil { + slog.Error("failed to start worker socket server", "error", err) + os.Exit(1) + } + slog.Info("worker ready", "redis_key_prefix", cfg.RedisKeyPrefix, "worker_queue", cfg.WorkerQueue, "handlers", runner.Len(), + "addr", fmt.Sprintf("%s:%d%s", cfg.Host, cfg.Port, cfg.SocketPath), ) + + <-ctx.Done() + if err := server.Stop(); err != nil { + slog.Error("failed to stop worker socket server", "error", err) + os.Exit(1) + } } diff --git a/services/worker/internal/config/config.go b/services/worker/internal/config/config.go index 64c72cf..dcd386c 100644 --- a/services/worker/internal/config/config.go +++ b/services/worker/internal/config/config.go @@ -7,6 +7,9 @@ type Config struct { RedisURL string RedisKeyPrefix string WorkerQueue string + Host string + Port int + SocketPath string } func Load() Config { @@ -15,6 +18,9 @@ func Load() Config { RedisURL: getenv("REDIS_URL", "redis://localhost:6379/0"), RedisKeyPrefix: getenv("ALT_REDIS_KEY_PREFIX", "alt"), WorkerQueue: getenv("ALT_WORKER_QUEUE", "default"), + Host: getenv("ALT_WORKER_HOST", "127.0.0.1"), + Port: getenvInt("ALT_WORKER_PORT", 8081), + SocketPath: getenv("ALT_WORKER_SOCKET_PATH", "/socket"), } } @@ -24,3 +30,15 @@ func getenv(key, fallback string) string { } return fallback } + +func getenvInt(key string, fallback int) int { + value := os.Getenv(key) + if value == "" { + return fallback + } + parsed, err := strconv.Atoi(value) + if err != nil { + return fallback + } + return parsed +} diff --git a/services/worker/internal/config/config_test.go b/services/worker/internal/config/config_test.go index db11eaf..297e8bf 100644 --- a/services/worker/internal/config/config_test.go +++ b/services/worker/internal/config/config_test.go @@ -2,76 +2,47 @@ package config import ( "os" + "strconv" "testing" ) -func TestLoadDefaults(t *testing.T) { - // Ensure env vars are clean for default test - envs := []string{"DATABASE_URL", "REDIS_URL", "ALT_REDIS_KEY_PREFIX", "ALT_WORKER_QUEUE"} - for _, env := range envs { - if val, ok := os.LookupEnv(env); ok { - defer os.Setenv(env, val) - os.Unsetenv(env) - } else { - defer os.Unsetenv(env) - } - } +func TestConfigLoad(t *testing.T) { + os.Setenv("ALT_WORKER_HOST", "127.0.0.9") + os.Setenv("ALT_WORKER_PORT", "9999") + os.Setenv("ALT_WORKER_SOCKET_PATH", "/ws-test") + defer func() { + os.Unsetenv("ALT_WORKER_HOST") + os.Unsetenv("ALT_WORKER_PORT") + os.Unsetenv("ALT_WORKER_SOCKET_PATH") + }() cfg := Load() - expectedDB := "postgres://alt:alt@localhost:5432/alt?sslmode=disable" - if cfg.DatabaseURL != expectedDB { - t.Errorf("expected DatabaseURL %q, got %q", expectedDB, cfg.DatabaseURL) + if cfg.Host != "127.0.0.9" { + t.Errorf("expected Host to be 127.0.0.9, got %q", cfg.Host) } - - expectedRedis := "redis://localhost:6379/0" - if cfg.RedisURL != expectedRedis { - t.Errorf("expected RedisURL %q, got %q", expectedRedis, cfg.RedisURL) + if cfg.Port != 9999 { + t.Errorf("expected Port to be 9999, got %d", cfg.Port) } - - expectedPrefix := "alt" - if cfg.RedisKeyPrefix != expectedPrefix { - t.Errorf("expected RedisKeyPrefix %q, got %q", expectedPrefix, cfg.RedisKeyPrefix) - } - - expectedQueue := "default" - if cfg.WorkerQueue != expectedQueue { - t.Errorf("expected WorkerQueue %q, got %q", expectedQueue, cfg.WorkerQueue) + if cfg.SocketPath != "/ws-test" { + t.Errorf("expected SocketPath to be /ws-test, got %q", cfg.SocketPath) } } -func TestLoadEnvOverrides(t *testing.T) { - envs := map[string]string{ - "DATABASE_URL": "postgres://user:pass@host:5432/db", - "REDIS_URL": "redis://redis-host:6379/1", - "ALT_REDIS_KEY_PREFIX": "custom-prefix", - "ALT_WORKER_QUEUE": "custom-queue", - } - - for k, v := range envs { - if val, ok := os.LookupEnv(k); ok { - defer os.Setenv(k, val) - } else { - defer os.Unsetenv(k) - } - os.Setenv(k, v) - } +func TestConfigDefault(t *testing.T) { + os.Unsetenv("ALT_WORKER_HOST") + os.Unsetenv("ALT_WORKER_PORT") + os.Unsetenv("ALT_WORKER_SOCKET_PATH") cfg := Load() - if cfg.DatabaseURL != envs["DATABASE_URL"] { - t.Errorf("expected DatabaseURL %q, got %q", envs["DATABASE_URL"], cfg.DatabaseURL) + if cfg.Host != "127.0.0.1" { + t.Errorf("expected default Host to be 127.0.0.1, got %q", cfg.Host) } - - if cfg.RedisURL != envs["REDIS_URL"] { - t.Errorf("expected RedisURL %q, got %q", envs["REDIS_URL"], cfg.RedisURL) + if cfg.Port != 8081 { + t.Errorf("expected default Port to be 8081, got %d", cfg.Port) } - - if cfg.RedisKeyPrefix != envs["ALT_REDIS_KEY_PREFIX"] { - t.Errorf("expected RedisKeyPrefix %q, got %q", envs["ALT_REDIS_KEY_PREFIX"], cfg.RedisKeyPrefix) - } - - if cfg.WorkerQueue != envs["ALT_WORKER_QUEUE"] { - t.Errorf("expected WorkerQueue %q, got %q", envs["ALT_WORKER_QUEUE"], cfg.WorkerQueue) + if cfg.SocketPath != "/socket" { + t.Errorf("expected default SocketPath to be /socket, got %q", cfg.SocketPath) } } diff --git a/services/worker/internal/contracts/parser_map.go b/services/worker/internal/contracts/parser_map.go new file mode 100644 index 0000000..5bbb7e7 --- /dev/null +++ b/services/worker/internal/contracts/parser_map.go @@ -0,0 +1,53 @@ +package contracts + +import ( + "google.golang.org/protobuf/proto" + + altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" + protoSocket "git.toki-labs.com/toki/proto-socket/go" +) + +func messageFactories() []func() proto.Message { + return []func() proto.Message{ + func() proto.Message { return &altv1.HelloRequest{} }, + func() proto.Message { return &altv1.HelloResponse{} }, + // market read surface: instruments and bars + func() proto.Message { return &altv1.ListInstrumentsRequest{} }, + func() proto.Message { return &altv1.ListInstrumentsResponse{} }, + func() proto.Message { return &altv1.ListBarsRequest{} }, + func() proto.Message { return &altv1.ListBarsResponse{} }, + // backtest surface: start / list / detail / result / compare + func() proto.Message { return &altv1.StartBacktestRequest{} }, + func() proto.Message { return &altv1.StartBacktestResponse{} }, + func() proto.Message { return &altv1.GetBacktestRunRequest{} }, + func() proto.Message { return &altv1.GetBacktestRunResponse{} }, + func() proto.Message { return &altv1.GetBacktestResultRequest{} }, + func() proto.Message { return &altv1.GetBacktestResultResponse{} }, + func() proto.Message { return &altv1.BacktestResult{} }, + func() proto.Message { return &altv1.ListBacktestRunsRequest{} }, + func() proto.Message { return &altv1.ListBacktestRunsResponse{} }, + func() proto.Message { return &altv1.GetBacktestRunDetailRequest{} }, + func() proto.Message { return &altv1.GetBacktestRunDetailResponse{} }, + func() proto.Message { return &altv1.CompareBacktestRunsRequest{} }, + func() proto.Message { return &altv1.CompareBacktestRunsResponse{} }, + } +} + +func ParserMap() protoSocket.ParserMap { + factories := messageFactories() + pm := make(protoSocket.ParserMap, len(factories)) + for _, factory := range factories { + pm[protoSocket.TypeNameOf(factory())] = parserFor(factory) + } + return pm +} + +func parserFor(factory func() proto.Message) func([]byte) (proto.Message, error) { + return func(data []byte) (proto.Message, error) { + msg := factory() + if err := proto.Unmarshal(data, msg); err != nil { + return nil, err + } + return msg, nil + } +} diff --git a/services/worker/internal/socket/handlers.go b/services/worker/internal/socket/handlers.go new file mode 100644 index 0000000..3fb8f20 --- /dev/null +++ b/services/worker/internal/socket/handlers.go @@ -0,0 +1,56 @@ +package socket + +import ( + protoSocket "git.toki-labs.com/toki/proto-socket/go" + + altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" +) + +type sessionHandler struct { + requestType string + register func(*protoSocket.WsClient) +} + +func sessionHandlers() []sessionHandler { + return []sessionHandler{ + helloHandler(), + } +} + +func registerSessionHandlers(client *protoSocket.WsClient) { + registerHandlers(client, sessionHandlers()) +} + +func registerHandlers(client *protoSocket.WsClient, handlers []sessionHandler) { + for _, handler := range handlers { + if handler.register == nil { + continue + } + handler.register(client) + } +} + +func helloHandler() sessionHandler { + return sessionHandler{ + requestType: protoSocket.TypeNameOf(&altv1.HelloRequest{}), + register: func(client *protoSocket.WsClient) { + protoSocket.AddRequestListenerTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, handleHello) + }, + } +} + +func handleHello(req *altv1.HelloRequest) (*altv1.HelloResponse, error) { + protocolVersion := req.GetAltProtocolVersion() + if protocolVersion == "" { + protocolVersion = defaultAltProtocolVersion + } + return &altv1.HelloResponse{ + ServerName: serverName, + ServerVersion: serverVersion, + AltProtocolVersion: protocolVersion, + Capabilities: []string{ + "hello", + "worker-execution", + }, + }, nil +} diff --git a/services/worker/internal/socket/server.go b/services/worker/internal/socket/server.go new file mode 100644 index 0000000..880fb04 --- /dev/null +++ b/services/worker/internal/socket/server.go @@ -0,0 +1,39 @@ +package socket + +import ( + "context" + + protoSocket "git.toki-labs.com/toki/proto-socket/go" + "nhooyr.io/websocket" + + "git.toki-labs.com/toki/alt/services/worker/internal/config" + workerContracts "git.toki-labs.com/toki/alt/services/worker/internal/contracts" +) + +const ( + serverName = "alt-worker" + serverVersion = "dev" + defaultAltProtocolVersion = "alt.v1" +) + +type Server struct { + wsServer *protoSocket.WsServer +} + +func NewServer(cfg config.Config) *Server { + wsServer := protoSocket.NewWsServer(cfg.Host, cfg.Port, cfg.SocketPath, func(conn *websocket.Conn) *protoSocket.WsClient { + return protoSocket.NewWsClient(conn, 30, 10, workerContracts.ParserMap()) + }) + wsServer.OnClientConnected = registerSessionHandlers + return &Server{ + wsServer: wsServer, + } +} + +func (s *Server) Start(ctx context.Context) error { + return s.wsServer.Start(ctx) +} + +func (s *Server) Stop() error { + return s.wsServer.Stop() +} diff --git a/services/worker/internal/socket/server_test.go b/services/worker/internal/socket/server_test.go new file mode 100644 index 0000000..7cc0208 --- /dev/null +++ b/services/worker/internal/socket/server_test.go @@ -0,0 +1,65 @@ +package socket + +import ( + "context" + "net" + "testing" + "time" + + altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" + "git.toki-labs.com/toki/alt/services/worker/internal/config" + workerContracts "git.toki-labs.com/toki/alt/services/worker/internal/contracts" + protoSocket "git.toki-labs.com/toki/proto-socket/go" +) + +func TestWorkerSocketServerHello(t *testing.T) { + // Get temporary port + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("failed to listen on temporary port: %v", err) + } + port := l.Addr().(*net.TCPAddr).Port + l.Close() + + cfg := config.Config{ + Host: "127.0.0.1", + Port: port, + SocketPath: "/socket", + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + server := NewServer(cfg) + if err := server.Start(ctx); err != nil { + t.Fatalf("failed to start worker socket server: %v", err) + } + defer server.Stop() + + // Wait slightly for server setup + time.Sleep(50 * time.Millisecond) + + // Dial + client, err := protoSocket.DialWs(ctx, "127.0.0.1", port, "/socket", workerContracts.ParserMap()) + if err != nil { + t.Fatalf("failed to dial worker socket server: %v", err) + } + defer client.Close() + + // Send HelloRequest + req := &altv1.HelloRequest{ + AltProtocolVersion: "alt.v1", + } + + res, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, req, 2*time.Second) + if err != nil { + t.Fatalf("failed to send HelloRequest: %v", err) + } + + if res.ServerName != "alt-worker" { + t.Errorf("expected ServerName to be %q, got %q", "alt-worker", res.ServerName) + } + if res.AltProtocolVersion != "alt.v1" { + t.Errorf("expected AltProtocolVersion to be %q, got %q", "alt.v1", res.AltProtocolVersion) + } +}