From 145363f6e4cf587ca5dba297728af51e358fd0b5 Mon Sep 17 00:00:00 2001 From: Natalie Diaz Date: Thu, 30 Jul 2026 11:08:26 -0700 Subject: [PATCH 1/2] Write usage logs when diverting to spanner --- internal/server/handler_v2.go | 98 +++++++++-------- internal/server/handler_v2_test.go | 102 ++++++++++++------ internal/server/remote/client.go | 45 ++++---- internal/server/remote/client_test.go | 97 +++++++++++++++++ internal/server/remote/datasource.go | 22 ++-- internal/server/v2/observation/observation.go | 72 +++++++++++-- .../server/v2/observation/observation_test.go | 94 ++++++++++++++++ 7 files changed, 420 insertions(+), 110 deletions(-) create mode 100644 internal/server/remote/client_test.go create mode 100644 internal/server/v2/observation/observation_test.go diff --git a/internal/server/handler_v2.go b/internal/server/handler_v2.go index 9a1374f0f..0902d1c4a 100644 --- a/internal/server/handler_v2.go +++ b/internal/server/handler_v2.go @@ -32,6 +32,7 @@ import ( "github.com/datacommonsorg/mixer/internal/server/translator" v2observation "github.com/datacommonsorg/mixer/internal/server/v2/observation" "github.com/datacommonsorg/mixer/internal/server/v2/resolve" + "github.com/datacommonsorg/mixer/internal/server/v2/shared" "github.com/datacommonsorg/mixer/internal/util" "github.com/google/uuid" "golang.org/x/sync/errgroup" @@ -369,54 +370,65 @@ func (s *Server) handleV2Event( func (s *Server) V2Observation( ctx context.Context, in *pbv2.ObservationRequest, ) (*pbv2.ObservationResponse, error) { + var v2Resp *pbv2.ObservationResponse + var queryType shared.QueryType + var err error + if s.shouldDivertV2(ctx) { - return s.dispatcher.Observation(ctx, in) + v2Resp, err = s.dispatcher.Observation(ctx, in) + if err != nil { + return nil, err + } + queryType = v2observation.GetQueryType(in) + } else { + v2StartTime := time.Now() + + surface, _ := util.GetMetadata(ctx) + + initialResp, qt, err := v2observation.ObservationInternal( + ctx, + s.store, + s.cachedata.Load(), + s.metadata, + s.httpClient, + in, + surface) + if err != nil { + return nil, err + } + queryType = qt + calculatedResps, err := v2observation.MaybeCalculateHoles( + ctx, + s.store, + s.cachedata.Load(), + s.metadata, + s.httpClient, + in, + initialResp, + surface, + ) + if err != nil { + return nil, err + } + // initialResp is preferred over any calculated response. + combinedResp := append([]*pbv2.ObservationResponse{initialResp}, calculatedResps...) + v2Resp = merger.MergeMultiObservation(combinedResp) + v2Latency := time.Since(v2StartTime) + + s.maybeMirrorV3( + ctx, + in, + v2Resp, + v2Latency, + func(ctx context.Context, req proto.Message) (proto.Message, error) { + return s.V3Observation(ctx, req.(*pbv2.ObservationRequest)) + }, + GetV2ObservationCmpOpts(), + ) } - v2StartTime := time.Now() - surface, toRemote := util.GetMetadata(ctx) - initialResp, queryType, err := v2observation.ObservationInternal( - ctx, - s.store, - s.cachedata.Load(), - s.metadata, - s.httpClient, - in, - surface) - if err != nil { - return nil, err - } - calculatedResps, err := v2observation.MaybeCalculateHoles( - ctx, - s.store, - s.cachedata.Load(), - s.metadata, - s.httpClient, - in, - initialResp, - surface, - ) - if err != nil { - return nil, err - } - // initialResp is preferred over any calculated response. - combinedResp := append([]*pbv2.ObservationResponse{initialResp}, calculatedResps...) - v2Resp := merger.MergeMultiObservation(combinedResp) - v2Latency := time.Since(v2StartTime) - - s.maybeMirrorV3( - ctx, - in, - v2Resp, - v2Latency, - func(ctx context.Context, req proto.Message) (proto.Message, error) { - return s.V3Observation(ctx, req.(*pbv2.ObservationRequest)) - }, - GetV2ObservationCmpOpts(), - ) - // Create a new ID to return as a header on the response. // This is used for usage logging and in the website to log cached usage. responseId := uuid.New() diff --git a/internal/server/handler_v2_test.go b/internal/server/handler_v2_test.go index e9308db89..2765c3632 100644 --- a/internal/server/handler_v2_test.go +++ b/internal/server/handler_v2_test.go @@ -26,6 +26,9 @@ import ( "github.com/datacommonsorg/mixer/internal/featureflags" pbv2 "github.com/datacommonsorg/mixer/internal/proto/v2" "github.com/datacommonsorg/mixer/internal/server/cache" + "github.com/datacommonsorg/mixer/internal/server/datasource" + "github.com/datacommonsorg/mixer/internal/server/datasources" + "github.com/datacommonsorg/mixer/internal/server/dispatcher" "github.com/datacommonsorg/mixer/internal/server/resource" v2observation "github.com/datacommonsorg/mixer/internal/server/v2/observation" "github.com/datacommonsorg/mixer/internal/server/v2/resolve" @@ -204,45 +207,82 @@ func TestObservationInternal(t *testing.T) { } } +type mockObservationDataSource struct { + datasource.DataSource +} + +func (m *mockObservationDataSource) Observation(ctx context.Context, req *pbv2.ObservationRequest) (*pbv2.ObservationResponse, error) { + return &pbv2.ObservationResponse{ + ByVariable: map[string]*pbv2.VariableObservation{ + "Count_Person": { + ByEntity: map[string]*pbv2.EntityObservation{ + "country/USA": {}, + }, + }, + }, + }, nil +} +func (m *mockObservationDataSource) ID() string { return "mock" } + func TestV2Observation_UsageLog(t *testing.T) { - ctx := metadata.NewIncomingContext(context.Background(), metadata.MD{}) - s := &Server{ - store: &store.Store{}, - metadata: &resource.Metadata{}, - flags: &featureflags.Flags{}, - writeUsageLogs: true, - } - s.cachedata.Store(&cache.Cache{}) - req := &pbv2.ObservationRequest{ - Select: []string{"variable", "entity", "date", "value"}, - Variable: &pbv2.DcidOrExpression{ - Dcids: []string{"Count_Person"}, + tests := []struct { + name string + useSpannerGraph bool + }{ + { + name: "legacy", + useSpannerGraph: false, }, - Entity: &pbv2.DcidOrExpression{ - Dcids: []string{"country/USA"}, + { + name: "diverted to dispatcher", + useSpannerGraph: true, }, } - // Capture slog output - var buf bytes.Buffer - handler := slog.NewTextHandler(&buf, nil) - logger := slog.New(handler) - originalLogger := slog.Default() - slog.SetDefault(logger) - defer slog.SetDefault(originalLogger) + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := metadata.NewIncomingContext(context.Background(), metadata.MD{}) + s := &Server{ + store: &store.Store{}, + metadata: &resource.Metadata{}, + flags: &featureflags.Flags{}, + writeUsageLogs: true, + useSpannerGraph: tc.useSpannerGraph, + dispatcher: dispatcher.NewDispatcher(nil, datasources.NewDataSources([]datasource.DataSource{&mockObservationDataSource{}}, nil)), + } + s.cachedata.Store(&cache.Cache{}) + req := &pbv2.ObservationRequest{ + Select: []string{"variable", "entity", "date", "value"}, + Variable: &pbv2.DcidOrExpression{ + Dcids: []string{"Count_Person"}, + }, + Entity: &pbv2.DcidOrExpression{ + Dcids: []string{"country/USA"}, + }, + } - _, _ = s.V2Observation(ctx, req) + // Capture slog output + var buf bytes.Buffer + handler := slog.NewTextHandler(&buf, nil) + logger := slog.New(handler) + originalLogger := slog.Default() + slog.SetDefault(logger) + defer slog.SetDefault(originalLogger) - outStr := strings.TrimSpace(buf.String()) + _, _ = s.V2Observation(ctx, req) - // Use regex to match the log message, ignoring the timestamp and pointer address. - wantLogRegex := `time=\S+ level=INFO msg=new_query usage_log.feature="{IsRemote:false Surface:}" usage_log.place_types=\[\] usage_log.query_type=value usage_log.stat_vars=\[0x[0-9a-f]+\] usage_log.response_id=\S+` - matched, err := regexp.MatchString(wantLogRegex, outStr) - if err != nil { - t.Fatalf("Failed to compile regex: %v", err) - } - if !matched { - t.Errorf("log output did not match expected pattern.\nGot: %s\nWant regex: %s", outStr, wantLogRegex) + outStr := strings.TrimSpace(buf.String()) + + // Use regex to match the log message, ignoring the timestamp and pointer address. + wantLogRegex := `time=\S+ level=INFO msg=new_query usage_log.feature="{IsRemote:false Surface:}" usage_log.place_types=\[\] usage_log.query_type=value usage_log.stat_vars=\[0x[0-9a-f]+\] usage_log.response_id=\S+` + matched, err := regexp.MatchString(wantLogRegex, outStr) + if err != nil { + t.Fatalf("Failed to compile regex: %v", err) + } + if !matched { + t.Errorf("log output did not match expected pattern.\nGot: %s\nWant regex: %s", outStr, wantLogRegex) + } + }) } } diff --git a/internal/server/remote/client.go b/internal/server/remote/client.go index 6c4d6f0e7..0d22b90d6 100644 --- a/internal/server/remote/client.go +++ b/internal/server/remote/client.go @@ -16,6 +16,7 @@ package remote import ( + "context" "fmt" "log/slog" "net/http" @@ -47,14 +48,22 @@ func NewRemoteClient(metadata *resource.Metadata) (*RemoteClient, error) { }, nil } -func (rc *RemoteClient) Node(req *pbv2.NodeRequest) (*pbv2.NodeResponse, error) { +func getSurface(ctx context.Context) string { + if ctx == nil { + return "" + } + surface, _ := util.GetMetadata(ctx) + return surface +} + +func (rc *RemoteClient) Node(ctx context.Context, req *pbv2.NodeRequest) (*pbv2.NodeResponse, error) { err := updateNodeRequestNextToken(req, rc.id) if err != nil { return nil, err } resp := &pbv2.NodeResponse{} - err = util.FetchRemote(rc.metadata, rc.httpClient, "/v2/node", req, resp) + err = util.FetchRemote(rc.metadata, rc.httpClient, "/v2/node", req, resp, getSurface(ctx)) if err != nil { return nil, err } @@ -67,72 +76,72 @@ func (rc *RemoteClient) Node(req *pbv2.NodeRequest) (*pbv2.NodeResponse, error) return resp, nil } -func (rc *RemoteClient) Observation(req *pbv2.ObservationRequest) (*pbv2.ObservationResponse, error) { +func (rc *RemoteClient) Observation(ctx context.Context, req *pbv2.ObservationRequest) (*pbv2.ObservationResponse, error) { resp := &pbv2.ObservationResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/observation", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/observation", req, resp, getSurface(ctx)) if err != nil { return nil, err } return resp, nil } -func (rc *RemoteClient) NodeSearch(req *pbv2.NodeSearchRequest) (*pbv2.NodeSearchResponse, error) { +func (rc *RemoteClient) NodeSearch(ctx context.Context, req *pbv2.NodeSearchRequest) (*pbv2.NodeSearchResponse, error) { resp := &pbv2.NodeSearchResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v3/node_search", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v3/node_search", req, resp, getSurface(ctx)) if err != nil { return nil, err } return resp, nil } -func (rc *RemoteClient) Resolve(req *pbv2.ResolveRequest) (*pbv2.ResolveResponse, error) { +func (rc *RemoteClient) Resolve(ctx context.Context, req *pbv2.ResolveRequest) (*pbv2.ResolveResponse, error) { resp := &pbv2.ResolveResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/resolve", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/resolve", req, resp, getSurface(ctx)) if err != nil { return nil, err } return resp, nil } -func (rc *RemoteClient) Sparql(req *pb.SparqlRequest) (*pb.QueryResponse, error) { +func (rc *RemoteClient) Sparql(ctx context.Context, req *pb.SparqlRequest) (*pb.QueryResponse, error) { resp := &pb.QueryResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/sparql", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/sparql", req, resp, getSurface(ctx)) if err != nil { return nil, err } return resp, nil } -func (rc *RemoteClient) Event(req *pbv2.EventRequest) (*pbv2.EventResponse, error) { +func (rc *RemoteClient) Event(ctx context.Context, req *pbv2.EventRequest) (*pbv2.EventResponse, error) { resp := &pbv2.EventResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/event", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/event", req, resp, getSurface(ctx)) if err != nil { return nil, err } return resp, nil } -func (rc *RemoteClient) BulkVariableInfo(req *pbv1.BulkVariableInfoRequest) (*pbv1.BulkVariableInfoResponse, error) { +func (rc *RemoteClient) BulkVariableInfo(ctx context.Context, req *pbv1.BulkVariableInfoRequest) (*pbv1.BulkVariableInfoResponse, error) { resp := &pbv1.BulkVariableInfoResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/bulk/info/variable", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/bulk/info/variable", req, resp, getSurface(ctx)) if err != nil { return nil, err } return resp, nil } -func (rc *RemoteClient) BulkVariableGroupInfo(req *pbv1.BulkVariableGroupInfoRequest) (*pbv1.BulkVariableGroupInfoResponse, error) { +func (rc *RemoteClient) BulkVariableGroupInfo(ctx context.Context, req *pbv1.BulkVariableGroupInfoRequest) (*pbv1.BulkVariableGroupInfoResponse, error) { resp := &pbv1.BulkVariableGroupInfoResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/bulk/info/variable-group", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/bulk/info/variable-group", req, resp, getSurface(ctx)) if err != nil { return nil, err } return resp, nil } -func (rc *RemoteClient) FilterStatVarsByEntity(req *pb.FilterStatVarsByEntityRequest) (*pb.FilterStatVarsByEntityResponse, error) { +func (rc *RemoteClient) FilterStatVarsByEntity(ctx context.Context, req *pb.FilterStatVarsByEntityRequest) (*pb.FilterStatVarsByEntityResponse, error) { resp := &pb.FilterStatVarsByEntityResponse{} - err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/variable/filter", req, resp) + err := util.FetchRemote(rc.metadata, rc.httpClient, "/v2/variable/filter", req, resp, getSurface(ctx)) if err != nil { slog.Error("Failed to fetch remote variable filter", "error", err) return nil, err diff --git a/internal/server/remote/client_test.go b/internal/server/remote/client_test.go new file mode 100644 index 000000000..181052e25 --- /dev/null +++ b/internal/server/remote/client_test.go @@ -0,0 +1,97 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package remote + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + pbv2 "github.com/datacommonsorg/mixer/internal/proto/v2" + "github.com/datacommonsorg/mixer/internal/server/resource" + "google.golang.org/grpc/metadata" +) + +func TestGetSurface(t *testing.T) { + tests := []struct { + name string + ctx context.Context + want string + }{ + { + name: "nil context", + ctx: nil, + want: "", + }, + { + name: "empty background context", + ctx: context.Background(), + want: "", + }, + { + name: "context with x-surface header", + ctx: metadata.NewIncomingContext(context.Background(), metadata.Pairs("x-surface", "mcp-server")), + want: "mcp-server", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := getSurface(tc.ctx); got != tc.want { + t.Errorf("getSurface(%v) = %q, want %q", tc.ctx, got, tc.want) + } + }) + } +} + +func TestRemoteClient_Observation_SurfaceHeader(t *testing.T) { + var receivedSurface string + var receivedRemote string + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + receivedSurface = r.Header.Get("X-Surface") + receivedRemote = r.Header.Get("X-Remote") + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{}`)) + })) + defer ts.Close() + + meta := &resource.Metadata{ + RemoteMixerDomain: ts.URL, + RemoteMixerAPIKey: "test-api-key", + } + + client, err := NewRemoteClient(meta) + if err != nil { + t.Fatalf("NewRemoteClient failed: %v", err) + } + + rds := NewRemoteDataSource(client) + + ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs("x-surface", "test-surface-agent")) + _, err = rds.Observation(ctx, &pbv2.ObservationRequest{}) + if err != nil { + t.Fatalf("rds.Observation failed: %v", err) + } + + if receivedSurface != "test-surface-agent" { + t.Errorf("expected X-Surface header %q, got %q", "test-surface-agent", receivedSurface) + } + if receivedRemote != "true" { + t.Errorf("expected X-Remote header %q, got %q", "true", receivedRemote) + } +} diff --git a/internal/server/remote/datasource.go b/internal/server/remote/datasource.go index ba9fd9bfd..efbfe9e88 100644 --- a/internal/server/remote/datasource.go +++ b/internal/server/remote/datasource.go @@ -46,37 +46,37 @@ func (rds *RemoteDataSource) Id() string { func (rds *RemoteDataSource) Node(ctx context.Context, req *pbv2.NodeRequest, pageSize int) (*pbv2.NodeResponse, error) { // The remote datasource currently calls V2 node, which does not use custom pageSize. - // TODO: Propagate ctx through RemoteClient and FetchRemote and configure a - // bounded HTTP timeout so remote requests honor cancellation and deadlines. - return rds.client.Node(req) + // TODO: Configure a bounded HTTP timeout in FetchRemote so remote requests honor + // cancellation and deadlines. + return rds.client.Node(ctx, req) } func (rds *RemoteDataSource) Observation(ctx context.Context, req *pbv2.ObservationRequest) (*pbv2.ObservationResponse, error) { - return rds.client.Observation(req) + return rds.client.Observation(ctx, req) } func (rds *RemoteDataSource) NodeSearch(ctx context.Context, req *pbv2.NodeSearchRequest) (*pbv2.NodeSearchResponse, error) { - return rds.client.NodeSearch(req) + return rds.client.NodeSearch(ctx, req) } func (rds *RemoteDataSource) Resolve(ctx context.Context, req *pbv2.ResolveRequest) (*pbv2.ResolveResponse, error) { - return rds.client.Resolve(req) + return rds.client.Resolve(ctx, req) } func (rds *RemoteDataSource) Sparql(ctx context.Context, req *pb.SparqlRequest) (*pb.QueryResponse, error) { - return rds.client.Sparql(req) + return rds.client.Sparql(ctx, req) } func (rds *RemoteDataSource) Event(ctx context.Context, req *pbv2.EventRequest) (*pbv2.EventResponse, error) { - return rds.client.Event(req) + return rds.client.Event(ctx, req) } func (rds *RemoteDataSource) BulkVariableInfo(ctx context.Context, req *pbv1.BulkVariableInfoRequest) (*pbv1.BulkVariableInfoResponse, error) { - return rds.client.BulkVariableInfo(req) + return rds.client.BulkVariableInfo(ctx, req) } func (rds *RemoteDataSource) BulkVariableGroupInfo(ctx context.Context, req *pbv1.BulkVariableGroupInfoRequest) (*pbv1.BulkVariableGroupInfoResponse, error) { - return rds.client.BulkVariableGroupInfo(req) + return rds.client.BulkVariableGroupInfo(ctx, req) } func (rds *RemoteDataSource) SdmxData(ctx context.Context, req *sdmxpb.SdmxDataQuery) (*sdmxpb.SdmxDataResult, error) { @@ -90,5 +90,5 @@ func (rds *RemoteDataSource) SdmxAvailability(ctx context.Context, req *sdmxpb.S } func (rds *RemoteDataSource) FilterStatVarsByEntity(ctx context.Context, req *pb.FilterStatVarsByEntityRequest) (*pb.FilterStatVarsByEntityResponse, error) { - return rds.client.FilterStatVarsByEntity(req) + return rds.client.FilterStatVarsByEntity(ctx, req) } diff --git a/internal/server/v2/observation/observation.go b/internal/server/v2/observation/observation.go index c5b2d8e72..87aa76589 100644 --- a/internal/server/v2/observation/observation.go +++ b/internal/server/v2/observation/observation.go @@ -33,6 +33,62 @@ import ( "google.golang.org/grpc/status" ) +// GetQueryType determines the QueryType from an ObservationRequest. +func GetQueryType(in *pbv2.ObservationRequest) shared.QueryType { + var queryDate, queryValue, queryVariable, queryEntity, queryFacet bool + for _, item := range in.GetSelect() { + switch item { + case "date": + queryDate = true + case "value": + queryValue = true + case "variable": + queryVariable = true + case "entity": + queryEntity = true + case "facet": + queryFacet = true + } + } + if !queryVariable || !queryEntity { + return "" + } + + variable := in.GetVariable() + entity := in.GetEntity() + + // Observation date and value query. + if queryDate && queryValue { + // Series or Collection. + if (len(variable.GetDcids()) > 0 && len(entity.GetDcids()) > 0) || + (len(variable.GetDcids()) > 0 && entity.GetExpression() != "") { + return shared.QueryTypeValue + } + // Derived series. + if variable.GetFormula() != "" && len(entity.GetDcids()) > 0 { + return shared.QueryTypeDerived + } + } + + // Get facet information for pair. + if !queryDate && !queryValue && queryFacet { + // Series or Collection. + if (len(variable.GetDcids()) > 0 && len(entity.GetDcids()) > 0) || + (len(variable.GetDcids()) > 0 && entity.GetExpression() != "") { + return shared.QueryTypeFacet + } + } + + // Get existence of pair. + if !queryDate && !queryValue { + if len(entity.GetDcids()) > 0 { + return shared.QueryTypeExistence + } + } + + return "" +} + func ObservationCore( ctx context.Context, store *store.Store, @@ -65,6 +121,8 @@ func ObservationCore( variable := in.GetVariable() entity := in.GetEntity() + queryType := GetQueryType(in) + // Observation date and value query. if queryDate && queryValue { // Series. @@ -79,7 +137,7 @@ func ObservationCore( in.GetFilter(), ) - return result, shared.QueryTypeValue, err + return result, queryType, err } // Collection. @@ -105,7 +163,7 @@ func ObservationCore( in.GetDate(), in.GetFilter(), ) - return res, shared.QueryTypeValue, err + return res, queryType, err } // Derived series. @@ -117,7 +175,7 @@ func ObservationCore( entity.GetDcids(), ) - return res, shared.QueryTypeDerived, err + return res, queryType, err } } @@ -133,7 +191,7 @@ func ObservationCore( entity.GetDcids(), ) - return res, shared.QueryTypeFacet, err + return res, queryType, err } // Collection if len(variable.GetDcids()) > 0 && entity.GetExpression() != "" { @@ -156,7 +214,7 @@ func ObservationCore( in.GetDate(), ) - return res, shared.QueryTypeFacet, err + return res, queryType, err } } @@ -167,13 +225,13 @@ func ObservationCore( // Have both entity.dcids and variable.dcids. Check existence cache. res, err := Existence( ctx, store, cachedata, variable.GetDcids(), entity.GetDcids()) - return res, shared.QueryTypeExistence, err + return res, queryType, err } // TODO: Support appending entities from entity.expression // Only have entity.dcids, fetch variables for each entity. res, err := Variable(ctx, store, entity.GetDcids()) - return res, shared.QueryTypeExistence, err + return res, queryType, err } } return &pbv2.ObservationResponse{}, "", nil diff --git a/internal/server/v2/observation/observation_test.go b/internal/server/v2/observation/observation_test.go new file mode 100644 index 000000000..b2a497912 --- /dev/null +++ b/internal/server/v2/observation/observation_test.go @@ -0,0 +1,94 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package observation + +import ( + "testing" + + pbv2 "github.com/datacommonsorg/mixer/internal/proto/v2" + "github.com/datacommonsorg/mixer/internal/server/v2/shared" +) + +func TestGetQueryType(t *testing.T) { + tests := []struct { + name string + in *pbv2.ObservationRequest + want shared.QueryType + }{ + { + name: "value query - dcids", + in: &pbv2.ObservationRequest{ + Select: []string{"date", "value", "variable", "entity"}, + Variable: &pbv2.DcidOrExpression{Dcids: []string{"Count_Person"}}, + Entity: &pbv2.DcidOrExpression{Dcids: []string{"country/USA"}}, + }, + want: shared.QueryTypeValue, + }, + { + name: "value query - expression", + in: &pbv2.ObservationRequest{ + Select: []string{"date", "value", "variable", "entity"}, + Variable: &pbv2.DcidOrExpression{Dcids: []string{"Count_Person"}}, + Entity: &pbv2.DcidOrExpression{Expression: "geoId/06<-containedInPlace+{typeOf: City}"}, + }, + want: shared.QueryTypeValue, + }, + { + name: "derived series", + in: &pbv2.ObservationRequest{ + Select: []string{"date", "value", "variable", "entity"}, + Variable: &pbv2.DcidOrExpression{Formula: "foo / bar"}, + Entity: &pbv2.DcidOrExpression{Dcids: []string{"country/USA"}}, + }, + want: shared.QueryTypeDerived, + }, + { + name: "facet query", + in: &pbv2.ObservationRequest{ + Select: []string{"facet", "variable", "entity"}, + Variable: &pbv2.DcidOrExpression{Dcids: []string{"Count_Person"}}, + Entity: &pbv2.DcidOrExpression{Dcids: []string{"country/USA"}}, + }, + want: shared.QueryTypeFacet, + }, + { + name: "existence query", + in: &pbv2.ObservationRequest{ + Select: []string{"variable", "entity"}, + Variable: &pbv2.DcidOrExpression{Dcids: []string{"Count_Person"}}, + Entity: &pbv2.DcidOrExpression{Dcids: []string{"country/USA"}}, + }, + want: shared.QueryTypeExistence, + }, + { + name: "missing variable select", + in: &pbv2.ObservationRequest{ + Select: []string{"entity", "date", "value"}, + Variable: &pbv2.DcidOrExpression{Dcids: []string{"Count_Person"}}, + Entity: &pbv2.DcidOrExpression{Dcids: []string{"country/USA"}}, + }, + want: "", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := GetQueryType(tc.in) + if got != tc.want { + t.Errorf("GetQueryType() = %v, want %v", got, tc.want) + } + }) + } +} From 994d955f7ba525acd79a406f82d75fa8ce726e15 Mon Sep 17 00:00:00 2001 From: Natalie Diaz Date: Thu, 30 Jul 2026 11:56:13 -0700 Subject: [PATCH 2/2] fix --- internal/server/handler_v2.go | 5 +---- internal/server/remote/client_test.go | 2 +- 2 files changed, 2 insertions(+), 5 deletions(-) diff --git a/internal/server/handler_v2.go b/internal/server/handler_v2.go index 0902d1c4a..4a78f02f2 100644 --- a/internal/server/handler_v2.go +++ b/internal/server/handler_v2.go @@ -373,6 +373,7 @@ func (s *Server) V2Observation( var v2Resp *pbv2.ObservationResponse var queryType shared.QueryType var err error + surface, toRemote := util.GetMetadata(ctx) if s.shouldDivertV2(ctx) { v2Resp, err = s.dispatcher.Observation(ctx, in) @@ -383,8 +384,6 @@ func (s *Server) V2Observation( } else { v2StartTime := time.Now() - surface, _ := util.GetMetadata(ctx) - initialResp, qt, err := v2observation.ObservationInternal( ctx, s.store, @@ -427,8 +426,6 @@ func (s *Server) V2Observation( ) } - surface, toRemote := util.GetMetadata(ctx) - // Create a new ID to return as a header on the response. // This is used for usage logging and in the website to log cached usage. responseId := uuid.New() diff --git a/internal/server/remote/client_test.go b/internal/server/remote/client_test.go index 181052e25..dd33ada4b 100644 --- a/internal/server/remote/client_test.go +++ b/internal/server/remote/client_test.go @@ -66,7 +66,7 @@ func TestRemoteClient_Observation_SurfaceHeader(t *testing.T) { receivedRemote = r.Header.Get("X-Remote") w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusOK) - w.Write([]byte(`{}`)) + _, _ = w.Write([]byte(`{}`)) })) defer ts.Close()