mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: implement filters for the organizations query (#14468)
Required for organization sync. Allows fetching a filtered set of orgs.
This commit is contained in:
@@ -176,7 +176,7 @@ func (r *RootCmd) newCreateAdminUserCommand() *serpent.Command {
|
||||
// Create the user.
|
||||
var newUser database.User
|
||||
err = db.InTx(func(tx database.Store) error {
|
||||
orgs, err := tx.GetOrganizations(ctx)
|
||||
orgs, err := tx.GetOrganizations(ctx, database.GetOrganizationsParams{})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get organizations: %w", err)
|
||||
}
|
||||
|
||||
@@ -60,7 +60,7 @@ func TestServerCreateAdminUser(t *testing.T) {
|
||||
require.EqualValues(t, []string{codersdk.RoleOwner}, user.RBACRoles, "user does not have owner role")
|
||||
|
||||
// Check that user is admin in every org.
|
||||
orgs, err := db.GetOrganizations(ctx)
|
||||
orgs, err := db.GetOrganizations(ctx, database.GetOrganizationsParams{})
|
||||
require.NoError(t, err)
|
||||
orgIDs := make(map[uuid.UUID]struct{}, len(orgs))
|
||||
for _, org := range orgs {
|
||||
|
||||
@@ -1700,9 +1700,9 @@ func (q *querier) GetOrganizationIDsByMemberIDs(ctx context.Context, ids []uuid.
|
||||
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.GetOrganizationIDsByMemberIDs)(ctx, ids)
|
||||
}
|
||||
|
||||
func (q *querier) GetOrganizations(ctx context.Context) ([]database.Organization, error) {
|
||||
func (q *querier) GetOrganizations(ctx context.Context, args database.GetOrganizationsParams) ([]database.Organization, error) {
|
||||
fetch := func(ctx context.Context, _ interface{}) ([]database.Organization, error) {
|
||||
return q.db.GetOrganizations(ctx)
|
||||
return q.db.GetOrganizations(ctx, args)
|
||||
}
|
||||
return fetchWithPostFilter(q.auth, policy.ActionRead, fetch)(ctx, nil)
|
||||
}
|
||||
|
||||
@@ -635,7 +635,7 @@ func (s *MethodTestSuite) TestOrganization() {
|
||||
def, _ := db.GetDefaultOrganization(context.Background())
|
||||
a := dbgen.Organization(s.T(), db, database.Organization{})
|
||||
b := dbgen.Organization(s.T(), db, database.Organization{})
|
||||
check.Args().Asserts(def, policy.ActionRead, a, policy.ActionRead, b, policy.ActionRead).Returns(slice.New(def, a, b))
|
||||
check.Args(database.GetOrganizationsParams{}).Asserts(def, policy.ActionRead, a, policy.ActionRead, b, policy.ActionRead).Returns(slice.New(def, a, b))
|
||||
}))
|
||||
s.Run("GetOrganizationsByUserID", s.Subtest(func(db database.Store, check *expects) {
|
||||
u := dbgen.User(s.T(), db, database.User{})
|
||||
|
||||
@@ -3034,14 +3034,24 @@ func (q *FakeQuerier) GetOrganizationIDsByMemberIDs(_ context.Context, ids []uui
|
||||
return getOrganizationIDsByMemberIDRows, nil
|
||||
}
|
||||
|
||||
func (q *FakeQuerier) GetOrganizations(_ context.Context) ([]database.Organization, error) {
|
||||
func (q *FakeQuerier) GetOrganizations(_ context.Context, args database.GetOrganizationsParams) ([]database.Organization, error) {
|
||||
q.mutex.RLock()
|
||||
defer q.mutex.RUnlock()
|
||||
|
||||
if len(q.organizations) == 0 {
|
||||
return nil, sql.ErrNoRows
|
||||
tmp := make([]database.Organization, 0)
|
||||
for _, org := range q.organizations {
|
||||
if len(args.IDs) > 0 {
|
||||
if !slices.Contains(args.IDs, org.ID) {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if args.Name != "" && !strings.EqualFold(org.Name, args.Name) {
|
||||
continue
|
||||
}
|
||||
tmp = append(tmp, org)
|
||||
}
|
||||
return q.organizations, nil
|
||||
|
||||
return tmp, nil
|
||||
}
|
||||
|
||||
func (q *FakeQuerier) GetOrganizationsByUserID(_ context.Context, userID uuid.UUID) ([]database.Organization, error) {
|
||||
@@ -3060,9 +3070,7 @@ func (q *FakeQuerier) GetOrganizationsByUserID(_ context.Context, userID uuid.UU
|
||||
organizations = append(organizations, organization)
|
||||
}
|
||||
}
|
||||
if len(organizations) == 0 {
|
||||
return nil, sql.ErrNoRows
|
||||
}
|
||||
|
||||
return organizations, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -46,7 +46,7 @@ func TestInTx(t *testing.T) {
|
||||
go func() {
|
||||
<-inTx
|
||||
for i := 0; i < 20; i++ {
|
||||
orgs, err := uut.GetOrganizations(context.Background())
|
||||
orgs, err := uut.GetOrganizations(context.Background(), database.GetOrganizationsParams{})
|
||||
if err != nil {
|
||||
assert.ErrorIs(t, err, sql.ErrNoRows)
|
||||
}
|
||||
|
||||
@@ -865,9 +865,9 @@ func (m metricsStore) GetOrganizationIDsByMemberIDs(ctx context.Context, ids []u
|
||||
return organizations, err
|
||||
}
|
||||
|
||||
func (m metricsStore) GetOrganizations(ctx context.Context) ([]database.Organization, error) {
|
||||
func (m metricsStore) GetOrganizations(ctx context.Context, args database.GetOrganizationsParams) ([]database.Organization, error) {
|
||||
start := time.Now()
|
||||
organizations, err := m.s.GetOrganizations(ctx)
|
||||
organizations, err := m.s.GetOrganizations(ctx, args)
|
||||
m.queryLatencies.WithLabelValues("GetOrganizations").Observe(time.Since(start).Seconds())
|
||||
return organizations, err
|
||||
}
|
||||
|
||||
@@ -1750,18 +1750,18 @@ func (mr *MockStoreMockRecorder) GetOrganizationIDsByMemberIDs(arg0, arg1 any) *
|
||||
}
|
||||
|
||||
// GetOrganizations mocks base method.
|
||||
func (m *MockStore) GetOrganizations(arg0 context.Context) ([]database.Organization, error) {
|
||||
func (m *MockStore) GetOrganizations(arg0 context.Context, arg1 database.GetOrganizationsParams) ([]database.Organization, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetOrganizations", arg0)
|
||||
ret := m.ctrl.Call(m, "GetOrganizations", arg0, arg1)
|
||||
ret0, _ := ret[0].([]database.Organization)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetOrganizations indicates an expected call of GetOrganizations.
|
||||
func (mr *MockStoreMockRecorder) GetOrganizations(arg0 any) *gomock.Call {
|
||||
func (mr *MockStoreMockRecorder) GetOrganizations(arg0, arg1 any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetOrganizations", reflect.TypeOf((*MockStore)(nil).GetOrganizations), arg0)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetOrganizations", reflect.TypeOf((*MockStore)(nil).GetOrganizations), arg0, arg1)
|
||||
}
|
||||
|
||||
// GetOrganizationsByUserID mocks base method.
|
||||
|
||||
@@ -180,7 +180,7 @@ type sqlcQuerier interface {
|
||||
GetOrganizationByID(ctx context.Context, id uuid.UUID) (Organization, error)
|
||||
GetOrganizationByName(ctx context.Context, name string) (Organization, error)
|
||||
GetOrganizationIDsByMemberIDs(ctx context.Context, ids []uuid.UUID) ([]GetOrganizationIDsByMemberIDsRow, error)
|
||||
GetOrganizations(ctx context.Context) ([]Organization, error)
|
||||
GetOrganizations(ctx context.Context, arg GetOrganizationsParams) ([]Organization, error)
|
||||
GetOrganizationsByUserID(ctx context.Context, userID uuid.UUID) ([]Organization, error)
|
||||
GetParameterSchemasByJobID(ctx context.Context, jobID uuid.UUID) ([]ParameterSchema, error)
|
||||
GetPreviousTemplateVersion(ctx context.Context, arg GetPreviousTemplateVersionParams) (TemplateVersion, error)
|
||||
|
||||
@@ -516,7 +516,7 @@ func TestDefaultOrg(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Should start with the default org
|
||||
all, err := db.GetOrganizations(ctx)
|
||||
all, err := db.GetOrganizations(ctx, database.GetOrganizationsParams{})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, all, 1)
|
||||
require.True(t, all[0].IsDefault, "first org should always be default")
|
||||
@@ -1211,7 +1211,7 @@ func TestExpectOne(t *testing.T) {
|
||||
dbgen.Organization(t, db, database.Organization{})
|
||||
|
||||
// Organizations is an easy table without foreign key dependencies
|
||||
_, err = database.ExpectOne(db.GetOrganizations(ctx))
|
||||
_, err = database.ExpectOne(db.GetOrganizations(ctx, database.GetOrganizationsParams{}))
|
||||
require.ErrorContains(t, err, "too many rows returned")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -4596,10 +4596,28 @@ SELECT
|
||||
id, name, description, created_at, updated_at, is_default, display_name, icon
|
||||
FROM
|
||||
organizations
|
||||
WHERE
|
||||
true
|
||||
-- Filter by ids
|
||||
AND CASE
|
||||
WHEN array_length($1 :: uuid[], 1) > 0 THEN
|
||||
id = ANY($1)
|
||||
ELSE true
|
||||
END
|
||||
AND CASE
|
||||
WHEN $2::text != '' THEN
|
||||
LOWER("name") = LOWER($2)
|
||||
ELSE true
|
||||
END
|
||||
`
|
||||
|
||||
func (q *sqlQuerier) GetOrganizations(ctx context.Context) ([]Organization, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getOrganizations)
|
||||
type GetOrganizationsParams struct {
|
||||
IDs []uuid.UUID `db:"ids" json:"ids"`
|
||||
Name string `db:"name" json:"name"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) GetOrganizations(ctx context.Context, arg GetOrganizationsParams) ([]Organization, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getOrganizations, pq.Array(arg.IDs), arg.Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -12,7 +12,21 @@ LIMIT
|
||||
SELECT
|
||||
*
|
||||
FROM
|
||||
organizations;
|
||||
organizations
|
||||
WHERE
|
||||
true
|
||||
-- Filter by ids
|
||||
AND CASE
|
||||
WHEN array_length(@ids :: uuid[], 1) > 0 THEN
|
||||
id = ANY(@ids)
|
||||
ELSE true
|
||||
END
|
||||
AND CASE
|
||||
WHEN @name::text != '' THEN
|
||||
LOWER("name") = LOWER(@name)
|
||||
ELSE true
|
||||
END
|
||||
;
|
||||
|
||||
-- name: GetOrganizationByID :one
|
||||
SELECT
|
||||
|
||||
@@ -3,6 +3,7 @@ package coderd
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/httpapi"
|
||||
"github.com/coder/coder/v2/coderd/httpmw"
|
||||
@@ -18,7 +19,7 @@ import (
|
||||
// @Router /organizations [get]
|
||||
func (api *API) organizations(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
organizations, err := api.Database.GetOrganizations(ctx)
|
||||
organizations, err := api.Database.GetOrganizations(ctx, database.GetOrganizationsParams{})
|
||||
if httpapi.Is404Error(err) {
|
||||
httpapi.ResourceNotFound(rw)
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user