diff --git a/README.md b/README.md index f4dd6b0..d84ce92 100644 --- a/README.md +++ b/README.md @@ -37,7 +37,7 @@ To publish services from a regular YTsaurus operation, add the `task_proxy` anno }> ``` -`protocol` must be `http` or `grpc`. If `tasks_info` is omitted, task-proxy publishes every job port as an HTTP service named `port`. +`protocol` must be `http`, `grpc`, or `websocket`. A `websocket` service is proxied like `http`, but Envoy additionally accepts the `Upgrade: websocket` handshake on its routes; after the upgrade the connection is tunneled to the job. Note that `stream_idle_timeout_seconds` still applies to an idle WebSocket connection, so raise it or set it to `0` for connections that stay silent for a long time. If `tasks_info` is omitted, task-proxy publishes every job port as an HTTP service named `port`. The annotation can also override request timeouts for every service in that operation: diff --git a/server/pkg/discovery.go b/server/pkg/discovery.go index 4a5dff5..b3511ad 100644 --- a/server/pkg/discovery.go +++ b/server/pkg/discovery.go @@ -429,11 +429,12 @@ func parseTaskProxyAnnotation(taskProxyAny any) ([]taskServiceInfo, TaskTimeoutO if !ok { continue } - protocol, ok := protocolAny.(string) + protocolString, ok := protocolAny.(string) if !ok { continue } - if protocol != string(HTTP) && protocol != string(GRPC) { + protocol, ok := ParseProtocol(protocolString) + if !ok { continue } portIndexAny, ok := info[taskProxyPortIndexKey] @@ -451,7 +452,7 @@ func parseTaskProxyAnnotation(taskProxyAny any) ([]taskServiceInfo, TaskTimeoutO taskServiceInfos = append(taskServiceInfos, taskServiceInfo{ task: task, service: service, - protocol: Protocol(protocol), + protocol: protocol, portIndex: portIndex, }) } diff --git a/server/pkg/discovery_test.go b/server/pkg/discovery_test.go index b37fc83..010d7e0 100644 --- a/server/pkg/discovery_test.go +++ b/server/pkg/discovery_test.go @@ -37,6 +37,28 @@ func TestParseTaskProxyAnnotation(t *testing.T) { }, }, }, + { + name: "websocket protocol", + annotation: map[string]any{ + "enabled": true, + "tasks_info": map[string]any{ + "ui": map[string]any{ + "ws": map[string]any{ + "protocol": "websocket", + "port_index": 1, + }, + }, + }, + }, + expected: []taskServiceInfo{ + { + task: "ui", + service: "ws", + protocol: WEBSOCKET, + portIndex: 1, + }, + }, + }, { name: "minimal annotation", annotation: map[string]any{ diff --git a/server/pkg/task.go b/server/pkg/task.go index 5ce4411..51ec367 100644 --- a/server/pkg/task.go +++ b/server/pkg/task.go @@ -12,8 +12,18 @@ type Protocol string const ( HTTP Protocol = "http" GRPC Protocol = "grpc" + // WEBSOCKET is plain HTTP/1.1 upstream with the WebSocket upgrade enabled on the route. + WEBSOCKET Protocol = "websocket" ) +func ParseProtocol(value string) (Protocol, bool) { + switch Protocol(value) { + case HTTP, GRPC, WEBSOCKET: + return Protocol(value), true + } + return "", false +} + type HostPort struct { host string port uint32 @@ -52,6 +62,7 @@ func (t *Task) OperationAlias() string { func (t *Task) IDWithHostPort() string { sb := strings.Builder{} sb.WriteString(t.ID()) + sb.WriteString(string(t.protocol)) for _, job := range t.jobs { sb.WriteString(job.host) fmt.Fprintf(&sb, "%d", job.port) diff --git a/server/pkg/task_test.go b/server/pkg/task_test.go index 79b440f..34d5711 100644 --- a/server/pkg/task_test.go +++ b/server/pkg/task_test.go @@ -78,6 +78,24 @@ func TestValidateTask(t *testing.T) { } } +func TestTaskIDWithHostPortIncludesProtocol(t *testing.T) { + http := Task{operationID: "op", taskName: "task", service: "service", protocol: HTTP} + websocket := http + websocket.protocol = WEBSOCKET + + assert.NotEqual(t, http.IDWithHostPort(), websocket.IDWithHostPort()) +} + +func TestParseProtocol(t *testing.T) { + for _, value := range []string{"http", "grpc", "websocket"} { + protocol, ok := ParseProtocol(value) + assert.True(t, ok) + assert.Equal(t, Protocol(value), protocol) + } + _, ok := ParseProtocol("dns") + assert.False(t, ok) +} + func TestTaskIDWithHostPortIncludesTimeoutOverrides(t *testing.T) { withoutOverrides := Task{operationID: "op", taskName: "task", service: "service"} withOverride := withoutOverrides diff --git a/server/pkg/xds.go b/server/pkg/xds.go index 026956f..e4e7cd5 100644 --- a/server/pkg/xds.go +++ b/server/pkg/xds.go @@ -39,6 +39,8 @@ import ( const ( extAuthClusterName = "extAuthz" + websocketUpgradeType = "websocket" + idRouterHeaderName = "x-yt-taskproxy-id" // hash operationIDRouterHeaderName = "x-yt-taskproxy-operation-id" @@ -72,7 +74,7 @@ func makeSnapshot(hashToTask map[string]Task, version string, baseDomain string, var defaultVhostRoutes []*routev3.Route for hash, task := range hashToTask { - grpc := task.protocol == "grpc" + grpc := task.protocol == GRPC vhostName := fmt.Sprintf("%s-%s-%s", task.operationID, task.taskName, task.service) var vhostClusters []*routev3.WeightedCluster_ClusterWeight @@ -84,17 +86,24 @@ func makeSnapshot(hashToTask map[string]Task, version string, baseDomain string, Weight: &wrapperspb.UInt32Value{Value: 1}, }) } - action := &routev3.Route_Route{ - Route: &routev3.RouteAction{ - Timeout: durationpb.New(task.timeoutOverrides.routeTimeoutOr(timeoutConfig.RouteTimeout)), - IdleTimeout: durationpb.New(task.timeoutOverrides.streamIdleTimeoutOr(timeoutConfig.StreamIdleTimeout)), - ClusterSpecifier: &routev3.RouteAction_WeightedClusters{ - WeightedClusters: &routev3.WeightedCluster{ - Clusters: vhostClusters, - }, + routeAction := &routev3.RouteAction{ + Timeout: durationpb.New(task.timeoutOverrides.routeTimeoutOr(timeoutConfig.RouteTimeout)), + IdleTimeout: durationpb.New(task.timeoutOverrides.streamIdleTimeoutOr(timeoutConfig.StreamIdleTimeout)), + ClusterSpecifier: &routev3.RouteAction_WeightedClusters{ + WeightedClusters: &routev3.WeightedCluster{ + Clusters: vhostClusters, }, }, } + if task.protocol == WEBSOCKET { + // Enabling the upgrade on the route is enough: Envoy allows it even when the HCM lists no upgrade_configs. + // The upstream stays HTTP/1.1, which is what WebSocket needs. + routeAction.UpgradeConfigs = []*routev3.RouteAction_UpgradeConfig{{ + UpgradeType: websocketUpgradeType, + Enabled: wrapperspb.Bool(true), + }} + } + action := &routev3.Route_Route{Route: routeAction} domains := []string{getTaskHashDomain(hash, baseDomain)} if task.operationAlias != "" { domains = append(domains, getTaskAliasDomain(task, baseDomain)) diff --git a/server/pkg/xds_test.go b/server/pkg/xds_test.go index eb90f3c..7e26c06 100644 --- a/server/pkg/xds_test.go +++ b/server/pkg/xds_test.go @@ -284,6 +284,63 @@ func TestMakeSnapshotTimeouts(t *testing.T) { } } +func TestMakeSnapshotWebSocket(t *testing.T) { + hashToTask := map[string]Task{ + "ws123456": { + operationID: "op123", + operationAlias: "myalias", + taskName: "ui", + service: "ws", + protocol: WEBSOCKET, + jobs: []HostPort{{host: "10.0.0.1", port: 8080}}, + }, + "http1234": { + operationID: "op123", + taskName: "ui", + service: "api", + protocol: HTTP, + jobs: []HostPort{{host: "10.0.0.1", port: 8081}}, + }, + } + + snapshot, err := makeSnapshot(hashToTask, "v1", "example.com", false, true, DefaultTaskProxyTimeoutConfig()) + require.NoError(t, err) + + // WebSocket upstream must stay HTTP/1.1: no explicit HTTP/2 protocol options on the cluster. + cluster := snapshot.GetResources(resourcev3.ClusterType)["op123-ui-ws-0"].(*clusterv3.Cluster) + require.Empty(t, cluster.TypedExtensionProtocolOptions) + + listener := onlyListener(t, snapshot.GetResources(resourcev3.ListenerType)) + hcm := httpConnectionManager(t, listener) + require.Empty(t, hcm.UpgradeConfigs, "upgrade must be enabled per route, not globally") + + websocketRoutes, httpRoutes := 0, 0 + for _, vhost := range hcm.GetRouteConfig().GetVirtualHosts() { + for _, route := range vhost.GetRoutes() { + action := route.GetRoute() + if action == nil { + continue + } + switch action.GetWeightedClusters().GetClusters()[0].GetName() { + case "op123-ui-ws-0": + websocketRoutes++ + require.Len(t, action.UpgradeConfigs, 1) + require.Equal(t, "websocket", action.UpgradeConfigs[0].GetUpgradeType()) + require.True(t, action.UpgradeConfigs[0].GetEnabled().GetValue()) + case "op123-ui-api-0": + httpRoutes++ + require.Empty(t, action.UpgradeConfigs) + default: + t.Fatalf("unexpected cluster in route %v", route) + } + } + } + // domain vhost + id header + operation-id headers + alias headers + require.Equal(t, 4, websocketRoutes) + // domain vhost + id header + operation-id headers (no alias) + require.Equal(t, 3, httpRoutes) +} + func onlyListener(t *testing.T, resources map[string]cachetypes.Resource) *listenerv3.Listener { t.Helper() require.Len(t, resources, 1)