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
This commit is contained in:
toki 2026-05-30 19:14:07 +09:00
parent 07f970b030
commit 3ee268816b
18 changed files with 515 additions and 248 deletions

View file

@ -1,25 +0,0 @@
<!-- task=m-api-centered-proto-socket-rail/01_contracts_api_registry plan=0 tag=API -->
# CODE_REVIEW-cloud-G07: Contracts and API Handler Registry Foundation
## 구현 에이전트 소유 섹션
- 구현 요약: 미작성
- 변경 파일: 미작성
- 실행한 검증/명령: 미작성
- 남은 위험/후속 작업: 미작성
## 사용자 리뷰 요청
_기본값은 `없음`이다. 구현 중 사용자 결정, 외부 환경 준비, 또는 계획 범위 변경 없이는 안전하게 진행할 수 없으면 아래 항목을 실제 내용으로 교체하고, 구현을 중단한 뒤 active 파일을 그대로 둔 채 리뷰를 요청한다. code-review가 이 내용을 검증해 `USER_REVIEW.md`를 작성한다._
- 상태: 없음
- 사유 유형: 없음
- 결정 필요: 없음
- 차단 근거: 없음
- 실행한 검증/명령: 없음
- 재개 조건: 없음
## 코드 리뷰어 소유 섹션
- 리뷰 상태: 미작성
- 주요 발견사항: 미작성
- 테스트/검증 평가: 미작성
- 판정: 미작성

View file

@ -1,132 +0,0 @@
<!-- task=m-api-centered-proto-socket-rail/01_contracts_api_registry plan=0 tag=API -->
# 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`의 구현 에이전트 소유 섹션을 채운다. 이 파일 작성이 구현의 마지막 단계다.

View file

@ -19,6 +19,10 @@ func main() {
cfg := config.Load() cfg := config.Load()
server := socket.NewServer(cfg) server := socket.NewServer(cfg)
go func() {
_ = socket.Worker.Connect(ctx)
}()
if err := server.Start(ctx); err != nil { if err := server.Start(ctx); err != nil {
slog.Error("failed to start api socket server", "error", err) slog.Error("failed to start api socket server", "error", err)
os.Exit(1) os.Exit(1)
@ -30,4 +34,8 @@ func main() {
slog.Error("failed to stop api socket server", "error", err) slog.Error("failed to stop api socket server", "error", err)
os.Exit(1) os.Exit(1)
} }
if err := socket.Worker.Close(); err != nil {
slog.Error("failed to close worker socket client", "error", err)
}
} }

View file

@ -13,6 +13,7 @@ type Config struct {
HeartbeatIntervalSec int HeartbeatIntervalSec int
HeartbeatWaitSec int HeartbeatWaitSec int
WSOriginPatterns []string WSOriginPatterns []string
WorkerSocketURL string
} }
func Load() Config { func Load() Config {
@ -23,6 +24,7 @@ func Load() Config {
HeartbeatIntervalSec: getenvInt("ALT_SOCKET_HEARTBEAT_INTERVAL_SEC", 30), HeartbeatIntervalSec: getenvInt("ALT_SOCKET_HEARTBEAT_INTERVAL_SEC", 30),
HeartbeatWaitSec: getenvInt("ALT_SOCKET_HEARTBEAT_WAIT_SEC", 10), HeartbeatWaitSec: getenvInt("ALT_SOCKET_HEARTBEAT_WAIT_SEC", 10),
WSOriginPatterns: getenvList("ALT_API_WS_ORIGIN_PATTERNS", "localhost:*,127.0.0.1:*"), 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"),
} }
} }

View file

@ -31,3 +31,17 @@ func TestLoadParsesWebSocketOriginPatterns(t *testing.T) {
t.Fatalf("WSOriginPatterns = %v, want %v", cfg.WSOriginPatterns, want) 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)
}
}

View file

@ -7,29 +7,44 @@ import (
protoSocket "git.toki-labs.com/toki/proto-socket/go" 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. // ParserMap returns a fresh ParserMap populated with all ALT API messages.
func ParserMap() protoSocket.ParserMap { func ParserMap() protoSocket.ParserMap {
return protoSocket.ParserMap{ factories := messageFactories()
protoSocket.TypeNameOf(&altv1.HelloRequest{}): parserFor(func() proto.Message { return &altv1.HelloRequest{} }), pm := make(protoSocket.ParserMap, len(factories))
protoSocket.TypeNameOf(&altv1.HelloResponse{}): parserFor(func() proto.Message { return &altv1.HelloResponse{} }), for _, factory := range factories {
protoSocket.TypeNameOf(&altv1.ListInstrumentsRequest{}): parserFor(func() proto.Message { return &altv1.ListInstrumentsRequest{} }), pm[protoSocket.TypeNameOf(factory())] = parserFor(factory)
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{} }),
} }
return pm
} }
func parserFor(factory func() proto.Message) func([]byte) (proto.Message, error) { func parserFor(factory func() proto.Message) func([]byte) (proto.Message, error) {

View file

@ -9,16 +9,20 @@ import (
altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1"
) )
func TestParserMapIncludesAltMessages(t *testing.T) { // requiredAPIMessages is the milestone contract surface the API parser map must
pm := ParserMap() // 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
messages := []proto.Message{ // silently shrinking both sides together.
func requiredAPIMessages() []proto.Message {
return []proto.Message{
&altv1.HelloRequest{}, &altv1.HelloRequest{},
&altv1.HelloResponse{}, &altv1.HelloResponse{},
// market instruments / bars
&altv1.ListInstrumentsRequest{}, &altv1.ListInstrumentsRequest{},
&altv1.ListInstrumentsResponse{}, &altv1.ListInstrumentsResponse{},
&altv1.ListBarsRequest{}, &altv1.ListBarsRequest{},
&altv1.ListBarsResponse{}, &altv1.ListBarsResponse{},
// backtest start / list / detail / result / compare
&altv1.StartBacktestRequest{}, &altv1.StartBacktestRequest{},
&altv1.StartBacktestResponse{}, &altv1.StartBacktestResponse{},
&altv1.GetBacktestRunRequest{}, &altv1.GetBacktestRunRequest{},
@ -33,8 +37,12 @@ func TestParserMapIncludesAltMessages(t *testing.T) {
&altv1.CompareBacktestRunsRequest{}, &altv1.CompareBacktestRunsRequest{},
&altv1.CompareBacktestRunsResponse{}, &altv1.CompareBacktestRunsResponse{},
} }
}
for _, msg := range messages { func TestParserMapIncludesAltMessages(t *testing.T) {
pm := ParserMap()
for _, msg := range requiredAPIMessages() {
typeName := protoSocket.TypeNameOf(msg) typeName := protoSocket.TypeNameOf(msg)
t.Run(typeName, func(t *testing.T) { t.Run(typeName, func(t *testing.T) {
parser := pm[typeName] 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")
}
}

View file

@ -29,10 +29,16 @@ func sessionHandlers() []sessionHandler {
} }
// registerSessionHandlers wires every session handler onto a newly connected // registerSessionHandlers wires every session handler onto a newly connected
// client. nil registrars are skipped so a malformed registry entry cannot // client. It is the OnClientConnected entrypoint for the API socket server.
// panic the connection setup path.
func registerSessionHandlers(client *protoSocket.WsClient) { 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 { if handler.register == nil {
continue continue
} }

View file

@ -6,6 +6,7 @@ import (
"git.toki-labs.com/toki/alt/services/api/internal/config" "git.toki-labs.com/toki/alt/services/api/internal/config"
apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts" apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts"
"git.toki-labs.com/toki/alt/services/api/internal/workerclient"
) )
const ( const (
@ -14,7 +15,11 @@ const (
defaultAltProtocolVersion = "alt.v1" defaultAltProtocolVersion = "alt.v1"
) )
var Worker workerclient.WorkerClient
func NewServer(cfg config.Config) *protoSocket.WsServer { func NewServer(cfg config.Config) *protoSocket.WsServer {
Worker = workerclient.New(cfg.WorkerSocketURL)
options := protoSocket.WsServerOptions{} options := protoSocket.WsServerOptions{}
if len(cfg.WSOriginPatterns) > 0 { if len(cfg.WSOriginPatterns) > 0 {
options.AcceptOptions = &websocket.AcceptOptions{ options.AcceptOptions = &websocket.AcceptOptions{

View file

@ -95,20 +95,26 @@ func TestSessionHandlersCoverRequiredRequests(t *testing.T) {
} }
} }
func TestRegisterSessionHandlersSkipsNilRegistrar(t *testing.T) { func TestRegisterHandlersSkipsNilRegistrar(t *testing.T) {
// registerSessionHandlers must tolerate a malformed registry entry instead // registerHandlers must tolerate a malformed registry entry instead of
// of panicking during connection setup. // panicking during connection setup. A nil registrar is skipped without
// dereferencing the client.
defer func() { defer func() {
if r := recover(); r != nil { 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} called := false
if handler.register != nil { handlers := []sessionHandler{
t.Fatal("expected nil registrar for probe handler") {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 { func freeTCPPort(t *testing.T) int {

View file

@ -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
}

View file

@ -1,19 +1,42 @@
package main package main
import ( import (
"context"
"fmt"
"log/slog" "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/config"
"git.toki-labs.com/toki/alt/services/worker/internal/jobs" "git.toki-labs.com/toki/alt/services/worker/internal/jobs"
"git.toki-labs.com/toki/alt/services/worker/internal/socket"
) )
func main() { func main() {
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
cfg := config.Load() cfg := config.Load()
runner := jobs.NewRunner() runner := jobs.NewRunner()
jobs.RegisterBuiltins(runner) 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", slog.Info("worker ready",
"redis_key_prefix", cfg.RedisKeyPrefix, "redis_key_prefix", cfg.RedisKeyPrefix,
"worker_queue", cfg.WorkerQueue, "worker_queue", cfg.WorkerQueue,
"handlers", runner.Len(), "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)
}
} }

View file

@ -7,6 +7,9 @@ type Config struct {
RedisURL string RedisURL string
RedisKeyPrefix string RedisKeyPrefix string
WorkerQueue string WorkerQueue string
Host string
Port int
SocketPath string
} }
func Load() Config { func Load() Config {
@ -15,6 +18,9 @@ func Load() Config {
RedisURL: getenv("REDIS_URL", "redis://localhost:6379/0"), RedisURL: getenv("REDIS_URL", "redis://localhost:6379/0"),
RedisKeyPrefix: getenv("ALT_REDIS_KEY_PREFIX", "alt"), RedisKeyPrefix: getenv("ALT_REDIS_KEY_PREFIX", "alt"),
WorkerQueue: getenv("ALT_WORKER_QUEUE", "default"), 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 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
}

View file

@ -2,76 +2,47 @@ package config
import ( import (
"os" "os"
"strconv"
"testing" "testing"
) )
func TestLoadDefaults(t *testing.T) { func TestConfigLoad(t *testing.T) {
// Ensure env vars are clean for default test os.Setenv("ALT_WORKER_HOST", "127.0.0.9")
envs := []string{"DATABASE_URL", "REDIS_URL", "ALT_REDIS_KEY_PREFIX", "ALT_WORKER_QUEUE"} os.Setenv("ALT_WORKER_PORT", "9999")
for _, env := range envs { os.Setenv("ALT_WORKER_SOCKET_PATH", "/ws-test")
if val, ok := os.LookupEnv(env); ok { defer func() {
defer os.Setenv(env, val) os.Unsetenv("ALT_WORKER_HOST")
os.Unsetenv(env) os.Unsetenv("ALT_WORKER_PORT")
} else { os.Unsetenv("ALT_WORKER_SOCKET_PATH")
defer os.Unsetenv(env) }()
}
}
cfg := Load() cfg := Load()
expectedDB := "postgres://alt:alt@localhost:5432/alt?sslmode=disable" if cfg.Host != "127.0.0.9" {
if cfg.DatabaseURL != expectedDB { t.Errorf("expected Host to be 127.0.0.9, got %q", cfg.Host)
t.Errorf("expected DatabaseURL %q, got %q", expectedDB, cfg.DatabaseURL)
} }
if cfg.Port != 9999 {
expectedRedis := "redis://localhost:6379/0" t.Errorf("expected Port to be 9999, got %d", cfg.Port)
if cfg.RedisURL != expectedRedis {
t.Errorf("expected RedisURL %q, got %q", expectedRedis, cfg.RedisURL)
} }
if cfg.SocketPath != "/ws-test" {
expectedPrefix := "alt" t.Errorf("expected SocketPath to be /ws-test, got %q", cfg.SocketPath)
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)
} }
} }
func TestLoadEnvOverrides(t *testing.T) { func TestConfigDefault(t *testing.T) {
envs := map[string]string{ os.Unsetenv("ALT_WORKER_HOST")
"DATABASE_URL": "postgres://user:pass@host:5432/db", os.Unsetenv("ALT_WORKER_PORT")
"REDIS_URL": "redis://redis-host:6379/1", os.Unsetenv("ALT_WORKER_SOCKET_PATH")
"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)
}
cfg := Load() cfg := Load()
if cfg.DatabaseURL != envs["DATABASE_URL"] { if cfg.Host != "127.0.0.1" {
t.Errorf("expected DatabaseURL %q, got %q", envs["DATABASE_URL"], cfg.DatabaseURL) t.Errorf("expected default Host to be 127.0.0.1, got %q", cfg.Host)
} }
if cfg.Port != 8081 {
if cfg.RedisURL != envs["REDIS_URL"] { t.Errorf("expected default Port to be 8081, got %d", cfg.Port)
t.Errorf("expected RedisURL %q, got %q", envs["REDIS_URL"], cfg.RedisURL)
} }
if cfg.SocketPath != "/socket" {
if cfg.RedisKeyPrefix != envs["ALT_REDIS_KEY_PREFIX"] { t.Errorf("expected default SocketPath to be /socket, got %q", cfg.SocketPath)
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)
} }
} }

View file

@ -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
}
}

View file

@ -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
}

View file

@ -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()
}

View file

@ -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)
}
}