diff --git a/dax/controller/schemar/schemar_test.go b/dax/controller/schemar/schemar_test.go index 5649f61fc..e1501651e 100644 --- a/dax/controller/schemar/schemar_test.go +++ b/dax/controller/schemar/schemar_test.go @@ -301,12 +301,12 @@ func TestSQLSchemar(t *testing.T) { requireCode(t, err, cschemar.ErrCodeFieldNameInvalid) }) - t.Run("Test create table that already exists", func(t *testing.T) { + t.Run("Test create table where table name already exists", func(t *testing.T) { tx, err = trans.BeginTx(context.Background(), true) require.NoError(t, err) defer tx.Rollback() err = schemar.CreateTable(tx, qtbl) - requireCode(t, err, dax.ErrTableIDExists) + requireCode(t, err, dax.ErrTableNameExists) }) t.Run("Find database by name that doesn't exist", func(t *testing.T) { @@ -317,6 +317,17 @@ func TestSQLSchemar(t *testing.T) { requireCode(t, err, dax.ErrDatabaseNameDoesNotExist) }) + t.Run("Create database with database name that already exists", func(t *testing.T) { + tx, err = trans.BeginTx(context.Background(), true) + require.NoError(t, err) + defer tx.Rollback() + err = schemar.CreateDatabase(tx, + &dax.QualifiedDatabase{ + OrganizationID: orgID, + Database: dax.Database{ID: dbID2, Name: dbName}}) + requireCode(t, err, dax.ErrDatabaseNameExists) + }) + t.Run("Find database by ID that doesn't exist", func(t *testing.T) { tx, err = trans.BeginTx(context.Background(), true) require.NoError(t, err) @@ -335,6 +346,7 @@ func TestSQLSchemar(t *testing.T) { } func requireCode(t *testing.T, err error, code errors.Code) { + t.Helper() if !errors.Is(err, code) { t.Fatalf("Error '%v' does not have code %s.", err, code) } diff --git a/dax/controller/sqldb/schemar.go b/dax/controller/sqldb/schemar.go index 845ae95b9..0191b8b75 100644 --- a/dax/controller/sqldb/schemar.go +++ b/dax/controller/sqldb/schemar.go @@ -50,15 +50,23 @@ func (s *Schemar) CreateDatabase(tx dax.Transaction, qdb *dax.QualifiedDatabase) return dax.NewErrDatabaseIDExists(qdb.QualifiedID()) } + // Check if org exists, if not, create org. org := &models.Organization{ID: string(qdb.OrganizationID)} - if ok, err := dt.C.Where("id = ?", qdb.OrganizationID).Exists(org); err != nil { + if exists, err := dt.C.Where("id = ?", qdb.OrganizationID).Exists(org); err != nil { return errors.Wrap(err, "checking for org") - } else if !ok { + } else if !exists { if err := dt.C.Create(org); err != nil { return errors.Wrap(err, "creating organization") } } + // Check if database name exists in org, if does, throw error. + if exists, err := dt.C.Where("name = ? AND organization_id = ?", qdb.Name, org.ID).Exists(&models.Database{}); err != nil { + return errors.Wrap(err, "checking database name") + } else if exists { + return dax.NewErrDatabaseNameExists(qdb.Name) + } + db := toModelDatabase(qdb) if err := dt.C.Create(db); err != nil { @@ -244,6 +252,13 @@ func (s *Schemar) CreateTable(tx dax.Transaction, qtbl *dax.QualifiedTable) erro return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") } + // Check to see if table name exists for a database ID, and if so, throw error + if exists, err := dt.C.Where("name = ? AND database_id = ?", qtbl.Name, qtbl.DatabaseID).Exists(&models.Table{}); err != nil { + return errors.Wrap(err, "checking if table name exists") + } else if exists { + return dax.NewErrTableNameExists(qtbl.Name) + } + tbl := toModelTable(qtbl) err := dt.C.Eager().Create(tbl) diff --git a/dax/errors.go b/dax/errors.go index eb653fbb2..d573ec41c 100644 --- a/dax/errors.go +++ b/dax/errors.go @@ -12,6 +12,7 @@ const ( ErrDatabaseIDExists errors.Code = "DatabaseIDExists" ErrDatabaseIDDoesNotExist errors.Code = "DatabaseIDDoesNotExist" ErrDatabaseNameDoesNotExist errors.Code = "DatabaseNameDoesNotExist" + ErrDatabaseNameExists errors.Code = "DatabaseNameExists" ErrTableIDExists errors.Code = "TableIDExists" ErrTableKeyExists errors.Code = "TableKeyExists" @@ -59,6 +60,13 @@ func NewErrDatabaseNameDoesNotExist(dbName DatabaseName) error { ) } +func NewErrDatabaseNameExists(dbName DatabaseName) error { + return errors.New( + ErrDatabaseNameExists, + fmt.Sprintf("database name %s already exists", dbName), + ) +} + func NewErrTableIDDoesNotExist(qtid QualifiedTableID) error { return errors.New( ErrTableIDDoesNotExist,