Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 24 additions & 17 deletions pkg/tenancy/device.go
Original file line number Diff line number Diff line change
Expand Up @@ -93,10 +93,13 @@ func NewDeviceResolver(db *database.Database, collections Collections) *DeviceRe
return &DeviceResolver{db: db, collections: collections.withDefaults()}
}

// persistedOrganisationOwnership reads only the canonical owner. Unlike the
// resource collections, the organisation document never had a snake_case owner
// field: models.Organisation has always declared bson:"ownerId", and it is the
// only spelling the ownerId_1 index covers.
type persistedOrganisationOwnership struct {
Id primitive.ObjectID `bson:"_id"`
OwnerId primitive.ObjectID `bson:"ownerId"`
LegacyOwnerId string `bson:"owner_id"`
Id primitive.ObjectID `bson:"_id"`
OwnerId primitive.ObjectID `bson:"ownerId"`
}

type persistedUserOwnership struct {
Expand All @@ -118,7 +121,7 @@ func DeviceProjection() *database.Projection {
)
}

// LoadDevice reads the ownership fields of a source device by its key.
// LoadDevice reads the ownership fields of exactly one source device by its key.
func (r *DeviceResolver) LoadDevice(ctx context.Context, deviceKey string) (models.Device, error) {
if r == nil || r.db == nil || r.db.Client == nil {
return models.Device{}, fmt.Errorf("source device %q lookup requires a database", deviceKey)
Expand All @@ -127,17 +130,25 @@ func (r *DeviceResolver) LoadDevice(ctx context.Context, deviceKey string) (mode
databaseCtx, cancel := context.WithTimeout(ctx, r.db.Client.GetTimeout())
defer cancel()

var device models.Device
err := r.db.Client.FindOne(
var devices []models.Device
err := r.db.Client.Find(
databaseCtx,
r.collections.Database,
r.collections.Devices,
map[string]string{properties.DeviceKey: deviceKey},
DeviceProjection(),
).Into(&device)
options.Find().SetLimit(2),
).All(&devices)
if err != nil {
return models.Device{}, fmt.Errorf("find source device %q: %w", deviceKey, err)
}
if len(devices) == 0 {
return models.Device{}, fmt.Errorf("find source device %q: %w", deviceKey, mongo.ErrNoDocuments)
}
if len(devices) > 1 {
return models.Device{}, fmt.Errorf("source device key %q resolves to multiple persisted devices", deviceKey)
}
device := devices[0]
if device.Id.IsZero() {
return models.Device{}, fmt.Errorf("source device %q has no persisted identity", deviceKey)
}
Expand Down Expand Up @@ -294,10 +305,12 @@ func (r *DeviceResolver) findOrganisationsByOwner(ctx context.Context, ownerId p
databaseCtx, cancel := context.WithTimeout(ctx, r.db.Client.GetTimeout())
defer cancel()

filter := map[string]any{"$or": []any{
map[string]any{properties.OrganisationOwnerId: ownerId},
map[string]any{"owner_id": ownerId.Hex()},
}}
// A plain equality rather than an ownership $or: organisations only ever
// stored the canonical owner, so a second arm would add no reachable
// document while costing the ownerId_1 index. An $or is only index-served
// when every arm is index-bounded, and no owner_id index exists to bound
// one — the whole lookup would degrade to a collection scan.
filter := map[string]any{properties.OrganisationOwnerId: ownerId}
var organisations []persistedOrganisationOwnership
if err := r.db.Client.Find(databaseCtx, r.collections.Database, r.collections.Organisations, filter, options.Find().SetLimit(2)).All(&organisations); err != nil {
return nil, fmt.Errorf("find organisations owned by %s: %w", ownerId.Hex(), err)
Expand All @@ -306,12 +319,6 @@ func (r *DeviceResolver) findOrganisationsByOwner(ctx context.Context, ownerId p
if organisation.Id.IsZero() {
return nil, fmt.Errorf("organisation owned by %s has no persisted identity", ownerId.Hex())
}
if !organisation.OwnerId.IsZero() && organisation.OwnerId != ownerId {
return nil, fmt.Errorf("organisation %s has conflicting canonical owner %s and legacy owner %s", organisation.Id.Hex(), organisation.OwnerId.Hex(), ownerId.Hex())
}
if organisation.LegacyOwnerId != "" && organisation.LegacyOwnerId != ownerId.Hex() {
return nil, fmt.Errorf("organisation %s has conflicting legacy owner %s and resolved owner %s", organisation.Id.Hex(), organisation.LegacyOwnerId, ownerId.Hex())
}
}
return organisations, nil
}
93 changes: 71 additions & 22 deletions pkg/tenancy/device_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,16 @@ package tenancy

import (
"context"
"errors"
"reflect"
"testing"

"github.com/uug-ai/database/pkg/database"
"github.com/uug-ai/models/pkg/models"
"github.com/uug-ai/models/pkg/properties"
"go.mongodb.org/mongo-driver/bson/primitive"
"go.mongodb.org/mongo-driver/mongo"
"go.mongodb.org/mongo-driver/mongo/options"
)

func testResolver(collections Collections) (*DeviceResolver, *database.MockDatabase) {
Expand Down Expand Up @@ -245,18 +249,33 @@ func TestResolveDeviceRejectsInvalidMasterRelationship(t *testing.T) {
assertFailsClosed(t, resolver, models.Device{Key: "device-1", UserId: userId.Hex()}, "invalid master ownership")
}

func TestResolveDeviceRejectsConflictingOwnerFields(t *testing.T) {
// The organisation lookup must stay a plain canonical equality. An ownership
// $or here would be unreachable — organisations never stored a snake_case owner
// — and would cost the ownerId_1 index, because an $or is only index-served
// when every arm is index-bounded and no owner_id index exists.
func TestResolveDeviceLooksUpOrganisationOwnerByCanonicalFieldOnly(t *testing.T) {
resolver, mock := testResolver(Collections{})
ownerId := primitive.NewObjectID()
organisationId := primitive.NewObjectID()
mock.QueueFindOne(nil, mongo.ErrNoDocuments)
mock.QueueFindOne(map[string]any{"_id": ownerId}, nil)
mock.QueueFind([]persistedOrganisationOwnership{{
Id: primitive.NewObjectID(),
OwnerId: primitive.NewObjectID(),
LegacyOwnerId: ownerId.Hex(),
}}, nil)
mock.QueueFind([]persistedOrganisationOwnership{{Id: organisationId, OwnerId: ownerId}}, nil)

ownership, err := resolver.ResolveDevice(context.Background(), models.Device{Key: "device-1", UserId: ownerId.Hex()})
if err != nil {
t.Fatalf("resolve device: %v", err)
}
if ownership.OrganisationId != organisationId {
t.Fatalf("organisation = %s, want %s", ownership.OrganisationId.Hex(), organisationId.Hex())
}

assertFailsClosed(t, resolver, models.Device{Key: "device-1", UserId: ownerId.Hex()}, "conflicting canonical and legacy organisation owners")
if len(mock.FindCalls) != 1 {
t.Fatalf("Find calls = %d, want 1", len(mock.FindCalls))
}
want := map[string]any{properties.OrganisationOwnerId: ownerId}
if !reflect.DeepEqual(mock.FindCalls[0].Filter, want) {
t.Fatalf("owner filter = %#v, want %#v", mock.FindCalls[0].Filter, want)
}
}

func TestResolveDeviceRejectsMissingOrInvalidOwnership(t *testing.T) {
Expand Down Expand Up @@ -284,11 +303,11 @@ func TestLoadDeviceReadsConfiguredCollection(t *testing.T) {
resolver, mock := testResolver(Collections{Database: "Custom", Devices: "sources"})
deviceId := primitive.NewObjectID()
organisationId := primitive.NewObjectID()
mock.QueueFindOne(map[string]any{
"_id": deviceId,
"key": "device-1",
"organisationId": organisationId.Hex(),
}, nil)
mock.QueueFind([]models.Device{{
Id: deviceId,
Key: "device-1",
OrganisationId: organisationId.Hex(),
}}, nil)

device, err := resolver.LoadDevice(context.Background(), "device-1")
if err != nil {
Expand All @@ -297,32 +316,62 @@ func TestLoadDeviceReadsConfiguredCollection(t *testing.T) {
if device.Id != deviceId {
t.Fatalf("device id = %s, want %s", device.Id.Hex(), deviceId.Hex())
}
if len(mock.FindOneCalls) != 1 {
t.Fatalf("FindOne calls = %d, want 1", len(mock.FindOneCalls))
if len(mock.FindCalls) != 1 {
t.Fatalf("Find calls = %d, want 1", len(mock.FindCalls))
}
if mock.FindCalls[0].Db != "Custom" || mock.FindCalls[0].Collection != "sources" {
t.Fatalf("read %s/%s, want Custom/sources", mock.FindCalls[0].Db, mock.FindCalls[0].Collection)
}
if mock.FindOneCalls[0].Db != "Custom" || mock.FindOneCalls[0].Collection != "sources" {
t.Fatalf("read %s/%s, want Custom/sources", mock.FindOneCalls[0].Db, mock.FindOneCalls[0].Collection)
var findOptions *options.FindOptions
for _, option := range mock.FindCalls[0].Opts {
if typed, ok := option.(*options.FindOptions); ok {
findOptions = typed
}
}
if findOptions == nil || findOptions.Limit == nil || *findOptions.Limit != 2 {
t.Fatalf("find options = %#v, want limit 2", mock.FindCalls[0].Opts)
}
}

func TestLoadDeviceRejectsDocumentWithoutIdentity(t *testing.T) {
resolver, mock := testResolver(Collections{})
mock.QueueFindOne(map[string]any{"key": "device-1"}, nil)
mock.QueueFind([]models.Device{{Key: "device-1"}}, nil)

if _, err := resolver.LoadDevice(context.Background(), "device-1"); err == nil {
t.Fatal("a device document with no persisted identity must fail closed")
}
}

func TestLoadDeviceRejectsMissingDevice(t *testing.T) {
resolver, mock := testResolver(Collections{})
mock.QueueFind([]models.Device{}, nil)

if _, err := resolver.LoadDevice(context.Background(), "device-1"); !errors.Is(err, mongo.ErrNoDocuments) {
t.Fatalf("LoadDevice() error = %v, want mongo.ErrNoDocuments", err)
}
}

func TestLoadDeviceRejectsAmbiguousKey(t *testing.T) {
resolver, mock := testResolver(Collections{})
mock.QueueFind([]models.Device{
{Id: primitive.NewObjectID(), Key: "device-1"},
{Id: primitive.NewObjectID(), Key: "device-1"},
}, nil)

if _, err := resolver.LoadDevice(context.Background(), "device-1"); err == nil {
t.Fatal("a device key resolving to multiple documents must fail closed")
}
}

func TestResolveDeviceByKeyReturnsDeviceAndOwnership(t *testing.T) {
resolver, mock := testResolver(Collections{})
deviceId := primitive.NewObjectID()
organisationId := primitive.NewObjectID()
mock.QueueFindOne(map[string]any{
"_id": deviceId,
"key": "device-1",
"organisationId": organisationId.Hex(),
}, nil)
mock.QueueFind([]models.Device{{
Id: deviceId,
Key: "device-1",
OrganisationId: organisationId.Hex(),
}}, nil)

device, ownership, err := resolver.ResolveDeviceByKey(context.Background(), "device-1")
if err != nil {
Expand Down
Loading