diff --git a/v1/server/compile_handler.go b/v1/server/compile_handler.go index 05a9fa0e17..ab35ebdcea 100644 --- a/v1/server/compile_handler.go +++ b/v1/server/compile_handler.go @@ -95,42 +95,11 @@ type CompileFiltersRequestV1 struct { Metadata map[string]any `json:"-"` } -var compileFiltersKnownKeys = map[string]bool{ - "input": true, "query": true, "unknowns": true, "options": true, -} - func (r *CompileFiltersRequestV1) UnmarshalJSON(data []byte) error { - type Alias CompileFiltersRequestV1 - aux := &struct { - *Alias - }{ - Alias: (*Alias)(r), - } - - if err := util.UnmarshalJSON(data, aux); err != nil { - return err - } - - var raw map[string]json.RawMessage - if err := util.UnmarshalJSON(data, &raw); err != nil { - return err - } - - for key, val := range raw { - if compileFiltersKnownKeys[key] { - continue - } - if r.Metadata == nil { - r.Metadata = make(map[string]any) - } - var v any - if err := util.UnmarshalJSON(val, &v); err != nil { - return err - } - r.Metadata[key] = v - } - - return nil + type alias CompileFiltersRequestV1 + extra, err := types.UnmarshalExtras[CompileFiltersRequestV1](data, (*alias)(r)) + r.Metadata = extra + return err } type compileFiltersRequest struct { @@ -154,33 +123,14 @@ type CompileResponseV1 struct { Metadata map[string]any `json:"-"` } -var compileResponseReservedFields = map[string]bool{ - "result": true, "explanation": true, "metrics": true, "hints": true, +func (r CompileResponseV1) MarshalJSON() ([]byte, error) { + type alias CompileResponseV1 + return types.MarshalExtras[CompileResponseV1](alias(r), r.Metadata) } -func (r CompileResponseV1) MarshalJSON() ([]byte, error) { - type Alias CompileResponseV1 - data, err := json.Marshal(Alias(r)) - if err != nil { - return nil, err - } - - if len(r.Metadata) == 0 { - return data, nil - } - - var base map[string]any - if err := json.Unmarshal(data, &base); err != nil { - return nil, err - } - - for key, val := range r.Metadata { - if !compileResponseReservedFields[key] { - base[key] = val - } - } - - return json.Marshal(base) +func init() { + types.RegisterJSONFields[CompileFiltersRequestV1]() + types.RegisterJSONFields[CompileResponseV1]() } func (s *Server) v1CompileFilters(w http.ResponseWriter, r *http.Request) { diff --git a/v1/server/types/types.go b/v1/server/types/types.go index ddce176bcf..5dc145df2a 100644 --- a/v1/server/types/types.go +++ b/v1/server/types/types.go @@ -9,7 +9,7 @@ import ( "bytes" "encoding/json" "fmt" - "maps" + "reflect" "strings" "github.com/open-policy-agent/opa/v1/ast" @@ -17,6 +17,104 @@ import ( "github.com/open-policy-agent/opa/v1/util" ) +// jsonFields holds the set of JSON field names declared on registered types. +// It is populated at package init time (see RegisterJSONFields) and is +// read-only thereafter, so concurrent reads need no synchronization. +var jsonFields = map[reflect.Type]map[string]bool{} + +// RegisterJSONFields records the JSON field names declared on T so that +// UnmarshalExtras and MarshalExtras can distinguish known fields from extras. +// Call from an init() function in the package that defines T. +func RegisterJSONFields[T any]() { + t := reflect.TypeOf((*T)(nil)).Elem() + if t.Kind() != reflect.Struct { + panic(fmt.Sprintf("types: RegisterJSONFields[%s]: not a struct", t)) + } + f := make(map[string]bool, t.NumField()) + for i := range t.NumField() { + sf := t.Field(i) + tag := sf.Tag.Get("json") + if tag == "-" { + continue + } + name, _, _ := strings.Cut(tag, ",") + if name == "" { + name = sf.Name + } + f[name] = true + } + jsonFields[t] = f +} + +func knownJSONFields[T any]() map[string]bool { + t := reflect.TypeOf((*T)(nil)).Elem() + f, ok := jsonFields[t] + if !ok { + panic(fmt.Sprintf("types: %s not registered with RegisterJSONFields", t)) + } + return f +} + +// UnmarshalExtras decodes data into into (typically a *alias of T to avoid +// recursing through T's UnmarshalJSON), then returns any top-level JSON keys +// not declared on T. T must have been registered with RegisterJSONFields. +func UnmarshalExtras[T any](data []byte, into any) (map[string]any, error) { + if err := util.UnmarshalJSON(data, into); err != nil { + return nil, err + } + var raw map[string]json.RawMessage + if err := util.UnmarshalJSON(data, &raw); err != nil { + return nil, err + } + known := knownJSONFields[T]() + var extra map[string]any + for k, val := range raw { + if known[k] { + continue + } + var v any + if err := util.UnmarshalJSON(val, &v); err != nil { + return nil, err + } + if extra == nil { + extra = make(map[string]any) + } + extra[k] = v + } + return extra, nil +} + +// MarshalExtras marshals v (typically an alias of T to avoid recursing through +// T's MarshalJSON) and merges extra into the resulting JSON object, dropping +// any extra keys that collide with a JSON field declared on T. T must have +// been registered with RegisterJSONFields. +func MarshalExtras[T any](v any, extra map[string]any) ([]byte, error) { + data, err := json.Marshal(v) + if err != nil { + return nil, err + } + if len(extra) == 0 { + return data, nil + } + known := knownJSONFields[T]() + var base map[string]any + if err := json.Unmarshal(data, &base); err != nil { + return nil, err + } + for k, val := range extra { + if known[k] { + continue + } + base[k] = val + } + return json.Marshal(base) +} + +func init() { + RegisterJSONFields[DataRequestV1]() + RegisterJSONFields[DataResponseV1]() +} + // Error codes returned by OPA's REST API. const ( CodeInternal = "internal_error" @@ -154,65 +252,15 @@ type DataRequestV1 struct { } func (r *DataRequestV1) UnmarshalJSON(data []byte) error { - type Alias DataRequestV1 - aux := &struct { - *Alias - }{ - Alias: (*Alias)(r), - } - - var raw map[string]json.RawMessage - if err := util.UnmarshalJSON(data, &raw); err != nil { - return err - } - - if err := util.UnmarshalJSON(data, aux); err != nil { - return err - } - - r.Metadata = make(map[string]any) - for key, val := range raw { - if key != "input" { - var v any - if err := util.UnmarshalJSON(val, &v); err != nil { - return err - } - r.Metadata[key] = v - } - } - - if len(r.Metadata) == 0 { - r.Metadata = nil - } - - return nil + type alias DataRequestV1 + extra, err := UnmarshalExtras[DataRequestV1](data, (*alias)(r)) + r.Metadata = extra + return err } func (r DataRequestV1) MarshalJSON() ([]byte, error) { - type Alias DataRequestV1 - aux := struct { - *Alias - }{ - Alias: (*Alias)(&r), - } - - data, err := json.Marshal(aux) - if err != nil { - return nil, err - } - - if len(r.Metadata) == 0 { - return data, nil - } - - var base map[string]any - if err := json.Unmarshal(data, &base); err != nil { - return nil, err - } - - maps.Copy(base, r.Metadata) - - return json.Marshal(base) + type alias DataRequestV1 + return MarshalExtras[DataRequestV1](alias(r), r.Metadata) } // DataResponseV1 models the response message for Data API read operations. @@ -234,88 +282,15 @@ type DataResponseV1 struct { } func (r *DataResponseV1) UnmarshalJSON(data []byte) error { - type Alias DataResponseV1 - aux := &struct { - *Alias - }{ - Alias: (*Alias)(r), - } - - var raw map[string]json.RawMessage - if err := util.UnmarshalJSON(data, &raw); err != nil { - return err - } - - if err := util.UnmarshalJSON(data, aux); err != nil { - return err - } - - knownFields := map[string]bool{ - "decision_id": true, - "provenance": true, - "explanation": true, - "metrics": true, - "result": true, - "warning": true, - } - - r.Metadata = make(map[string]any) - for key, val := range raw { - if !knownFields[key] { - var v any - if err := util.UnmarshalJSON(val, &v); err != nil { - return err - } - r.Metadata[key] = v - } - } - - if len(r.Metadata) == 0 { - r.Metadata = nil - } - - return nil + type alias DataResponseV1 + extra, err := UnmarshalExtras[DataResponseV1](data, (*alias)(r)) + r.Metadata = extra + return err } func (r DataResponseV1) MarshalJSON() ([]byte, error) { - type Alias DataResponseV1 - aux := struct { - *Alias - }{ - Alias: (*Alias)(&r), - } - - data, err := json.Marshal(aux) - if err != nil { - return nil, err - } - - if len(r.Metadata) == 0 { - return data, nil - } - - // Reserved field names that must not be overridden by metadata. - reservedFields := map[string]bool{ - "decision_id": true, - "provenance": true, - "explanation": true, - "metrics": true, - "result": true, - "warning": true, - } - - var base map[string]any - if err := json.Unmarshal(data, &base); err != nil { - return nil, err - } - - for key, val := range r.Metadata { - if !reservedFields[key] { - base[key] = val - } - } - - return json.Marshal(base) + type alias DataResponseV1 + return MarshalExtras[DataResponseV1](alias(r), r.Metadata) } // Warning models DataResponse warnings diff --git a/v1/server/types/types_extensible_test.go b/v1/server/types/types_extensible_test.go index 06ca82f86d..22f57f8ff0 100644 --- a/v1/server/types/types_extensible_test.go +++ b/v1/server/types/types_extensible_test.go @@ -120,6 +120,60 @@ func TestDataResponseV1_ExtraFields(t *testing.T) { } } +func TestDataRequestV1_MarshalRoundTrip(t *testing.T) { + src := `{"input":{"user":"alice"},"trace_id":"abc","tenant":{"id":"t-1"}}` + + var req DataRequestV1 + if err := json.Unmarshal([]byte(src), &req); err != nil { + t.Fatalf("Unmarshal failed: %v", err) + } + + out, err := json.Marshal(req) + if err != nil { + t.Fatalf("Marshal failed: %v", err) + } + + var got map[string]any + if err := json.Unmarshal(out, &got); err != nil { + t.Fatalf("decode failed: %v", err) + } + + input, ok := got["input"].(map[string]any) + if !ok || input["user"] != "alice" { + t.Errorf("input not preserved: %v", got["input"]) + } + if got["trace_id"] != "abc" { + t.Errorf("extra trace_id not preserved: %v", got["trace_id"]) + } + tenant, ok := got["tenant"].(map[string]any) + if !ok || tenant["id"] != "t-1" { + t.Errorf("extra tenant not preserved: %v", got["tenant"]) + } +} + +func TestDataRequestV1_MarshalDoesNotOverrideInput(t *testing.T) { + inp := any(map[string]any{"user": "alice"}) + req := DataRequestV1{ + Input: &inp, + Metadata: map[string]any{"input": "should_be_ignored"}, + } + + out, err := json.Marshal(req) + if err != nil { + t.Fatalf("Marshal failed: %v", err) + } + + var got map[string]any + if err := json.Unmarshal(out, &got); err != nil { + t.Fatalf("decode failed: %v", err) + } + + input, ok := got["input"].(map[string]any) + if !ok || input["user"] != "alice" { + t.Errorf("input overridden by metadata: %v", got["input"]) + } +} + func TestDataResponseV1_RoundTrip(t *testing.T) { input := `{"result": {"allowed": true}, "decision_id": "xyz", "custom": "data", "metrics": {"timer_rego_query_eval_ns": 1000}}`