package control import ( "bytes" "context" "encoding/json" "io" "net/http" rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" "google.golang.org/protobuf/encoding/protojson" "google.golang.org/protobuf/proto" ) const ( jsonRPCVersion = "2.0" jsonRPCParseError = -32700 jsonRPCInvalid = -32600 jsonRPCMethodAbsent = -32601 jsonRPCServerError = -32000 defaultJSONRPCBody = 24 << 20 ) // JSONRPCHandler adapts unary Control methods to standard JSON-RPC 2.0. It // intentionally has no authentication in v1; callers should keep it loopback // bound unless they explicitly accept the debugging exposure. type JSONRPCHandler struct { Service *Service MaxBody int64 Decoder protojson.UnmarshalOptions Encoder protojson.MarshalOptions } func NewJSONRPCHandler(service *Service, maxBody int64) *JSONRPCHandler { if maxBody <= 0 { maxBody = defaultJSONRPCBody } return &JSONRPCHandler{Service: service, MaxBody: maxBody, Decoder: protojson.UnmarshalOptions{DiscardUnknown: false}, Encoder: protojson.MarshalOptions{UseProtoNames: false}} } func (handler *JSONRPCHandler) ServeHTTP(response http.ResponseWriter, request *http.Request) { if request.Method != http.MethodPost { response.Header().Set("Allow", http.MethodPost) http.Error(response, "method not allowed", http.StatusMethodNotAllowed) return } if handler == nil || handler.Service == nil { http.Error(response, "control service unavailable", http.StatusServiceUnavailable) return } body, err := io.ReadAll(io.LimitReader(request.Body, handler.MaxBody+1)) if err != nil { handler.writeRPCError(response, nil, jsonRPCParseError, "could not read request", nil) return } if int64(len(body)) > handler.MaxBody { handler.writeRPCError(response, nil, jsonRPCInvalid, "request body exceeds limit", nil) return } var envelope jsonRPCRequest decoder := json.NewDecoder(bytes.NewReader(body)) decoder.DisallowUnknownFields() if err := decoder.Decode(&envelope); err != nil { handler.writeRPCError(response, nil, jsonRPCParseError, "invalid JSON", nil) return } var trailing any if err := decoder.Decode(&trailing); err != io.EOF || envelope.JSONRPC != jsonRPCVersion || envelope.Method == "" || len(envelope.ID) == 0 || bytes.Equal(bytes.TrimSpace(envelope.ID), []byte("null")) { handler.writeRPCError(response, nil, jsonRPCInvalid, "invalid JSON-RPC request", nil) return } result, callErr := handler.call(request.Context(), envelope.Method, envelope.Params) if callErr != nil { code, message, data := jsonRPCDomainError(callErr) handler.writeRPCError(response, envelope.ID, code, message, data) return } handler.writeRPCResult(response, envelope.ID, result) } type jsonRPCRequest struct { JSONRPC string `json:"jsonrpc"` ID json.RawMessage `json:"id"` Method string `json:"method"` Params json.RawMessage `json:"params"` } type jsonRPCResponse struct { JSONRPC string `json:"jsonrpc"` ID json.RawMessage `json:"id"` Result json.RawMessage `json:"result,omitempty"` Error *jsonRPCError `json:"error,omitempty"` } type jsonRPCError struct { Code int `json:"code"` Message string `json:"message"` Data json.RawMessage `json:"data,omitempty"` } func (handler *JSONRPCHandler) call(ctx context.Context, method string, params json.RawMessage) (proto.Message, error) { decode := func(message proto.Message) error { if len(params) == 0 || bytes.Equal(bytes.TrimSpace(params), []byte("null")) { params = []byte("{}") } return handler.Decoder.Unmarshal(params, message) } switch method { case "listClients": input := &rvboxv1.ListClientsRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.ListClients(ctx, input) case "getClient": input := &rvboxv1.GetClientRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.GetClient(ctx, input) case "listCommands": input := &rvboxv1.ListCommandsRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.ListCommands(ctx, input) case "getCommand": input := &rvboxv1.GetCommandRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.GetCommand(ctx, input) case "runCommand": input := &rvboxv1.RunCommandRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.RunCommand(ctx, input) case "appendStdin": input := &rvboxv1.AppendStdinRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.AppendStdin(ctx, input) case "closeStdin": input := &rvboxv1.CloseStdinRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.CloseStdin(ctx, input) case "signalCommand": input := &rvboxv1.ControlSignalCommandRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.SignalCommand(ctx, input) case "getOutput": input := &rvboxv1.GetOutputRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.GetOutput(ctx, input) case "listStorageIncidents": input := &rvboxv1.ListStorageIncidentsRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.ListStorageIncidents(ctx, input) case "repairStorageIncident": input := &rvboxv1.RepairStorageIncidentRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.RepairStorageIncident(ctx, input) case "acknowledgeStorageIncident": input := &rvboxv1.AcknowledgeStorageIncidentRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.AcknowledgeStorageIncident(ctx, input) case "authorizeClientTakeover": input := &rvboxv1.AuthorizeClientTakeoverRequest{} if err := decode(input); err != nil { return nil, status.Error(codes.InvalidArgument, err.Error()) } return handler.Service.AuthorizeClientTakeover(ctx, input) default: return nil, status.Error(codes.Unimplemented, "method not found") } } func (handler *JSONRPCHandler) writeRPCResult(response http.ResponseWriter, id json.RawMessage, message proto.Message) { encoded, err := handler.Encoder.Marshal(message) if err != nil { handler.writeRPCError(response, id, jsonRPCServerError, "could not encode response", nil) return } response.Header().Set("Content-Type", "application/json") response.WriteHeader(http.StatusOK) _, _ = response.Write(mustJSON(jsonRPCResponse{JSONRPC: jsonRPCVersion, ID: append([]byte(nil), id...), Result: encoded})) } func (handler *JSONRPCHandler) writeRPCError(response http.ResponseWriter, id json.RawMessage, code int, message string, data json.RawMessage) { if len(id) == 0 { id = []byte("null") } response.Header().Set("Content-Type", "application/json") response.WriteHeader(http.StatusOK) _, _ = response.Write(mustJSON(jsonRPCResponse{JSONRPC: jsonRPCVersion, ID: append([]byte(nil), id...), Error: &jsonRPCError{Code: code, Message: message, Data: data}})) } func jsonRPCDomainError(err error) (int, string, json.RawMessage) { if err == nil { return 0, "", nil } converted := status.Convert(err) if converted.Code() == codes.Unimplemented && converted.Message() == "method not found" { return jsonRPCMethodAbsent, converted.Message(), nil } data := []byte(nil) for _, detail := range converted.Details() { if control, ok := detail.(*rvboxv1.ControlError); ok { data, _ = protojson.Marshal(control) break } } if converted.Code() == codes.InvalidArgument { return jsonRPCInvalid, converted.Message(), data } return jsonRPCServerError, converted.Message(), data } func mustJSON(value any) []byte { encoded, err := json.Marshal(value) if err != nil { return []byte(`{"jsonrpc":"2.0","id":null,"error":{"code":-32000,"message":"internal encoding error"}}`) } return encoded }