diff --git a/dax/controller/balancer.go b/dax/controller/balancer.go index 87edb3cfd..63a57fd26 100644 --- a/dax/controller/balancer.go +++ b/dax/controller/balancer.go @@ -13,8 +13,8 @@ type Balancer interface { // be either transferred to other workers or placed on the free job list. RemoveWorker(tx dax.Transaction, addr dax.Address) ([]dax.WorkerDiff, error) - // FreeWorkers dissociates the given workers from a database. - FreeWorkers(tx dax.Transaction, addrs ...dax.Address) error + // ReleaseWorkers dissociates the given workers from a database. + ReleaseWorkers(tx dax.Transaction, addrs ...dax.Address) error // AddJobs adds new jobs for the given database. AddJobs(tx dax.Transaction, roleType dax.RoleType, qtid dax.QualifiedTableID, jobs ...dax.Job) ([]dax.WorkerDiff, error) @@ -64,7 +64,7 @@ func (b *NopBalancer) AddWorker(tx dax.Transaction, node *dax.Node) ([]dax.Worke func (b *NopBalancer) RemoveWorker(tx dax.Transaction, addr dax.Address) ([]dax.WorkerDiff, error) { return []dax.WorkerDiff{}, nil } -func (b *NopBalancer) FreeWorkers(tx dax.Transaction, addrs ...dax.Address) error { +func (b *NopBalancer) ReleaseWorkers(tx dax.Transaction, addrs ...dax.Address) error { return nil } func (b *NopBalancer) AddJobs(tx dax.Transaction, roleType dax.RoleType, qtid dax.QualifiedTableID, jobs ...dax.Job) ([]dax.WorkerDiff, error) { diff --git a/dax/controller/balancer/balancer.go b/dax/controller/balancer/balancer.go index 6e7686f23..f93dad12f 100644 --- a/dax/controller/balancer/balancer.go +++ b/dax/controller/balancer/balancer.go @@ -28,7 +28,7 @@ type Balancer struct { // current represents the current state of worker/job assigments. current WorkerJobService - nodeService controller.NodeService + workerRegistry controller.WorkerRegistry // freeJobs is the set of jobs which have yet to be assigned to a worker. // This could be because there are no available workers, or because a worker @@ -44,47 +44,37 @@ type Balancer struct { } // New returns a new instance of Balancer. -func New(ns controller.NodeService, fjs FreeJobService, wjs WorkerJobService, fws FreeWorkerService, schemar schemar.Schemar, logger logger.Logger) *Balancer { +func New(wr controller.WorkerRegistry, fjs FreeJobService, wjs WorkerJobService, fws FreeWorkerService, schemar schemar.Schemar, logger logger.Logger) *Balancer { return &Balancer{ - current: wjs, - nodeService: ns, - freeJobs: fjs, - freeWorkers: fws, - schemar: schemar, - logger: logger, + current: wjs, + workerRegistry: wr, + freeJobs: fjs, + freeWorkers: fws, + schemar: schemar, + logger: logger, } } -// AddWorker adds the given Node to the Balancer's available worker pool. -// TODO(tlt): this method takes a Node (as opposed to a Worker) because in the -// future we may want to maintain separate worker pools based on RoleType -// (compute, translate, etc.). +// AddWorker adds the given Node to the Balancer's available worker pool. Note +// that a node is used for ALL of the role types specified. In other words, +// specifying roleTypes = {compute, translate}, does not mean that the node can +// be used as either a compute worker or a translate worker. It means that it +// will be used as both. func (b *Balancer) AddWorker(tx dax.Transaction, node *dax.Node) ([]dax.WorkerDiff, error) { - addr := node.Address - b.logger.Debugf("AddWorker(%s)", addr) + b.logger.Debugf("AddWorker(%s)", node.Address) - if err := b.nodeService.CreateNode(tx, addr, node); err != nil { - return nil, errors.Wrapf(err, "creating node on node service: %s", addr) + if err := b.workerRegistry.AddWorker(tx, node); err != nil { + return nil, errors.Wrapf(err, "creating node on node service: %s", node.Address) } diffs := NewInternalDiffs() - // This logic means that a node is used for ALL of the role types specified. - // In other words, specifying roleTypes = {compute, translate}, does not - // mean that the node can be used as either a compute worker or a translate - // worker. It means that it will be used as both. - for _, rt := range node.RoleTypes { - if err := b.addWorker(tx, rt, addr); err != nil { - return nil, errors.Wrapf(err, "adding worker: (%s) %s", rt, addr) - } - } - - // Process the freeWorkers. + // Process the newly added workers. // TODO(tlt): this is a little heavy-handed. I'm sure we'll need to be more - // intentional about knowing which databases needs workers, as opposed to + // intentional about knowing which databases need workers, as opposed to // this brute force loop over all databases every time. if diff, err := b.balance(tx); err != nil { - return nil, errors.Wrapf(err, "balancing new worker: %s", addr) + return nil, errors.Wrapf(err, "balancing new worker: %s", node.Address) } else { diffs.Merge(diff) } @@ -92,23 +82,9 @@ func (b *Balancer) AddWorker(tx dax.Transaction, node *dax.Node) ([]dax.WorkerDi return diffs.Output(), nil } -// addWorker adds a worker to the free worker list. From there, it can be used -// by any database which needs a worker. -func (b *Balancer) addWorker(tx dax.Transaction, roleType dax.RoleType, addr dax.Address) error { - // If this worker already exists, don't do anything. - if dbkey := b.current.DatabaseForWorker(tx, addr); dbkey != "" { - return nil - } +func (b *Balancer) assignMinWorkers(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID) (InternalDiffs, error) { + b.logger.Debugf("assigning min workers for '%s', '%s'", roleType, qdbid) - if err := b.freeWorkers.AddWorkers(tx, roleType, addr); err != nil { - return errors.Wrap(err, "adding free worker") - } - - return nil -} - -func (b *Balancer) assignMinWorkers(tx dax.Transaction, roleType dax.RoleType) (InternalDiffs, error) { - b.logger.Debugf("assigning min workers for '%s'", roleType) // Find out how many free workers we have. freeWorkers, err := b.freeWorkers.ListWorkers(tx, roleType) if err != nil { @@ -122,10 +98,12 @@ func (b *Balancer) assignMinWorkers(tx dax.Transaction, roleType dax.RoleType) ( return InternalDiffs{}, nil } - // Get all database and their minWorkerCount (Database.Options.WorkersMin). - qdbs, err := b.schemar.Databases(tx, "") + // Get database and its minWorkerCount (Database.Options.WorkersMin). This + // used to get all databases, but now this method is specific to a single + // database. That's why we just get the one here. + qdbs, err := b.schemar.Databases(tx, qdbid.OrganizationID, qdbid.DatabaseID) if err != nil { - return nil, errors.Wrap(err, "getting all database") + return nil, errors.Wrap(err, "getting database") } // Create a map[database]int where int is the number of workers required to @@ -244,13 +222,12 @@ func (b *Balancer) databaseHasJobs(tx dax.Transaction, roleType dax.RoleType, qd func (b *Balancer) RemoveWorker(tx dax.Transaction, addr dax.Address) ([]dax.WorkerDiff, error) { diffs := NewInternalDiffs() - ////// The rest is database specific. //////////// - - // See if the worker is assigned to a database. If it's not, return early. + // See if the worker is assigned to a database. If it is, disassociate the + // worker from all of its jobs for the database. dbkey := b.current.DatabaseForWorker(tx, addr) if dbkey != "" { qdbid := dbkey.QualifiedDatabaseID() - for _, rt := range []dax.RoleType{dax.RoleTypeCompute, dax.RoleTypeTranslate} { + for _, rt := range dax.AllRoleTypes { if diff, err := b.removeDatabaseWorker(tx, rt, qdbid, addr); err != nil { return nil, errors.Wrapf(err, "removing worker: (%s) %s", rt, addr) } else { @@ -259,15 +236,8 @@ func (b *Balancer) RemoveWorker(tx dax.Transaction, addr dax.Address) ([]dax.Wor } } - // Remove the worker from the free worker list (if it's there). - for _, rt := range []dax.RoleType{dax.RoleTypeCompute, dax.RoleTypeTranslate} { - if err := b.freeWorkers.RemoveWorker(tx, rt, addr); err != nil { - return nil, errors.Wrapf(err, "removing worker from free list: (%s) %s", rt, addr) - } - } - - // Remove the worker (i.e. Node) from the node service. - if err := b.nodeService.DeleteNode(tx, addr); err != nil { + // Remove the worker from the worker registry. + if err := b.workerRegistry.RemoveWorker(tx, addr); err != nil { return nil, errors.Wrapf(err, "deleting node from node service: %s", addr) } @@ -284,6 +254,8 @@ func (b *Balancer) RemoveWorker(tx dax.Transaction, addr dax.Address) ([]dax.Wor return diffs.Output(), nil } +// removeDatabaseWorker is used to remove a worker that has been associated with +// a database. The worker here is determined by address. func (b *Balancer) removeDatabaseWorker(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address) (InternalDiffs, error) { jobs, err := b.current.ListJobs(tx, roleType, qdbid, addr) if err != nil { @@ -291,13 +263,8 @@ func (b *Balancer) removeDatabaseWorker(tx dax.Transaction, roleType dax.RoleTyp } // Before removing the worker, mark its jobs as free. - if err := b.freeJobs.MergeJobs(tx, roleType, qdbid, jobs); err != nil { - return nil, errors.Wrap(err, "merging free jobs") - } - - // Remove the worker. - if err := b.current.DeleteWorker(tx, roleType, qdbid, addr); err != nil { - return nil, errors.Wrap(err, "deleting worker") + if err := b.freeJobs.MarkJobsAsFree(tx, roleType, qdbid, jobs); err != nil { + return nil, errors.Wrap(err, "marking jobs as free") } // Even though this may not be useful to the caller (for example, in the @@ -311,8 +278,8 @@ func (b *Balancer) removeDatabaseWorker(tx dax.Transaction, roleType dax.RoleTyp return diff, nil } -func (b *Balancer) FreeWorkers(tx dax.Transaction, addrs ...dax.Address) error { - return errors.Wrap(b.current.FreeWorkers(tx, addrs...), "freeing workers") +func (b *Balancer) ReleaseWorkers(tx dax.Transaction, addrs ...dax.Address) error { + return errors.Wrap(b.current.ReleaseWorkers(tx, addrs...), "freeing workers") } func (b *Balancer) AddJobs(tx dax.Transaction, roleType dax.RoleType, qtid dax.QualifiedTableID, jobs ...dax.Job) ([]dax.WorkerDiff, error) { @@ -396,6 +363,7 @@ func (b *Balancer) addDatabaseJobs(tx dax.Transaction, roleType dax.RoleType, qd if err != nil { return nil, errors.Wrapf(err, "getting workers jobs: %s", roleType) } + jset := dax.NewSet[dax.Job]() for _, workerInfo := range workerJobs { jset.Merge(dax.NewSet(workerInfo.Jobs...)) @@ -437,12 +405,12 @@ func (b *Balancer) addDatabaseJobs(tx dax.Transaction, roleType dax.RoleType, qd jobCounts[lowWorker]++ } - for worker, jobs := range jobsToCreate { - if err := b.current.CreateJobs(tx, roleType, qdbid, worker, jobs...); err != nil { + for addr, jobs := range jobsToCreate { + if err := b.current.AssignWorkerToJobs(tx, roleType, qdbid, addr, jobs...); err != nil { return nil, errors.Wrap(err, "creating job") } for _, job := range jobs { - diffs.Added(worker, job) + diffs.Added(addr, job) } } @@ -527,7 +495,7 @@ func (b *Balancer) BalanceDatabase(tx dax.Transaction, qdbid dax.QualifiedDataba func (b *Balancer) balanceDatabase(tx dax.Transaction, qdbid dax.QualifiedDatabaseID) (InternalDiffs, error) { diffs := NewInternalDiffs() - for _, role := range []dax.RoleType{dax.RoleTypeCompute, dax.RoleTypeTranslate} { + for _, role := range dax.AllRoleTypes { diff, err := b.balanceDatabaseForRole(tx, role, qdbid) if err != nil { return nil, errors.Wrapf(err, "getting worker count: (%s) %s", role, qdbid) @@ -544,8 +512,7 @@ func (b *Balancer) balanceDatabaseForRole(tx dax.Transaction, roleType dax.RoleT // Before balancing, make sure the database has its minimum number of // workers satisfied. - // TODO(tlt): make assignMinWorkers database specific. - if diff, err := b.assignMinWorkers(tx, roleType); err != nil { + if diff, err := b.assignMinWorkers(tx, roleType, qdbid); err != nil { return nil, errors.Wrapf(err, "assigning min workers: (%s) %s", roleType, qdbid) } else { diffs.Merge(diff) @@ -809,11 +776,11 @@ func (b *Balancer) workerForJob(tx dax.Transaction, roleType dax.RoleType, qdbid } func (b *Balancer) ReadNode(tx dax.Transaction, addr dax.Address) (*dax.Node, error) { - return b.nodeService.ReadNode(tx, addr) + return b.workerRegistry.Worker(tx, addr) } func (b *Balancer) Nodes(tx dax.Transaction) ([]*dax.Node, error) { - return b.nodeService.Nodes(tx) + return b.workerRegistry.Workers(tx) } type WorkerJobService interface { @@ -823,10 +790,9 @@ type WorkerJobService interface { ListWorkers(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID) (dax.Addresses, error) CreateWorker(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address) error - DeleteWorker(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address) error - FreeWorkers(tx dax.Transaction, addrs ...dax.Address) error + ReleaseWorkers(tx dax.Transaction, addrs ...dax.Address) error - CreateJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address, job ...dax.Job) error + AssignWorkerToJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address, job ...dax.Job) error DeleteJob(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address, job dax.Job) error DeleteJobsForTable(tx dax.Transaction, roleType dax.RoleType, qtid dax.QualifiedTableID) (InternalDiffs, error) JobCounts(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr ...dax.Address) (map[dax.Address]int, error) @@ -840,12 +806,10 @@ type FreeJobService interface { DeleteJob(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, job dax.Job) error DeleteJobsForTable(tx dax.Transaction, roleType dax.RoleType, qtid dax.QualifiedTableID) error ListJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID) (dax.Jobs, error) - MergeJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, jobs dax.Jobs) error + MarkJobsAsFree(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, jobs dax.Jobs) error } type FreeWorkerService interface { - AddWorkers(tx dax.Transaction, roleType dax.RoleType, addrs ...dax.Address) error - RemoveWorker(tx dax.Transaction, roleType dax.RoleType, addr dax.Address) error PopWorkers(tx dax.Transaction, roleType dax.RoleType, num int) ([]dax.Address, error) ListWorkers(tx dax.Transaction, roleType dax.RoleType) (dax.Addresses, error) } diff --git a/dax/controller/balancer/free_job_test.go b/dax/controller/balancer/free_job_test.go index dd72fb848..31477940c 100644 --- a/dax/controller/balancer/free_job_test.go +++ b/dax/controller/balancer/free_job_test.go @@ -39,6 +39,11 @@ func TestFreeJobService(t *testing.T) { job2 := dax.Job(qtid.Key() + "job2") job3 := dax.Job(qtid.Key() + "job3") + node := &dax.Node{ + Address: nodeAddr, + RoleTypes: []dax.RoleType{role}, + } + err = fjSvc.CreateJobs(tx, role, qdbid, job1, job2, job3) require.NoError(t, err) @@ -49,22 +54,22 @@ func TestFreeJobService(t *testing.T) { require.NoError(t, err) require.ElementsMatch(t, dax.Jobs{job1, job3}, jobs) - fwSvc := sqldb.NewFreeWorkerService(nil) - err = fwSvc.AddWorkers(tx, role, nodeAddr) + workerReg := sqldb.NewWorkerRegistry(nil) + err = workerReg.AddWorker(tx, node) require.NoError(t, err) wjSvc := sqldb.NewWorkerJobService(nil) err = wjSvc.CreateWorker(tx, role, qdbid, nodeAddr) require.NoError(t, err) - err = wjSvc.CreateJobs(tx, role, qdbid, nodeAddr, job1) + err = wjSvc.AssignWorkerToJobs(tx, role, qdbid, nodeAddr, job1) require.NoError(t, err) jobs, err = fjSvc.ListJobs(tx, role, qdbid) require.NoError(t, err) require.ElementsMatch(t, dax.Jobs{job3}, jobs) - err = fjSvc.MergeJobs(tx, role, qdbid, dax.Jobs{job1}) + err = fjSvc.MarkJobsAsFree(tx, role, qdbid, dax.Jobs{job1}) require.NoError(t, err) jobs, err = fjSvc.ListJobs(tx, role, qdbid) diff --git a/dax/controller/balancer/free_worker_test.go b/dax/controller/balancer/free_worker_test.go index 56dff94d6..7c0f11141 100644 --- a/dax/controller/balancer/free_worker_test.go +++ b/dax/controller/balancer/free_worker_test.go @@ -20,12 +20,25 @@ func TestFreeWorkerService(t *testing.T) { } }() - fwSvc := sqldb.NewFreeWorkerService(nil) - err = fwSvc.AddWorkers(tx, role, nodeAddr, nodeAddr2, nodeAddr3, nodeAddr4, nodeAddr5) - require.NoError(t, err) + node1 := &dax.Node{Address: nodeAddr, RoleTypes: dax.AllRoleTypes} + node2 := &dax.Node{Address: nodeAddr2, RoleTypes: dax.AllRoleTypes} + node3 := &dax.Node{Address: nodeAddr3, RoleTypes: dax.AllRoleTypes} + node4 := &dax.Node{Address: nodeAddr4, RoleTypes: dax.AllRoleTypes} + node5 := &dax.Node{Address: nodeAddr5, RoleTypes: dax.AllRoleTypes} - err = fwSvc.RemoveWorker(tx, role, nodeAddr2) - require.NoError(t, err) + workerReg := sqldb.NewWorkerRegistry(nil) + + // Add some workers. + require.NoError(t, workerReg.AddWorker(tx, node1)) + require.NoError(t, workerReg.AddWorker(tx, node2)) + require.NoError(t, workerReg.AddWorker(tx, node3)) + require.NoError(t, workerReg.AddWorker(tx, node4)) + require.NoError(t, workerReg.AddWorker(tx, node5)) + + // Remove one of the workers. + require.NoError(t, workerReg.RemoveWorker(tx, node2.Address)) + + fwSvc := sqldb.NewFreeWorkerService(nil) addrs, err := fwSvc.ListWorkers(tx, role) require.NoError(t, err) diff --git a/dax/controller/balancer/node_test.go b/dax/controller/balancer/node_test.go index 186b31bb3..cacd8d869 100644 --- a/dax/controller/balancer/node_test.go +++ b/dax/controller/balancer/node_test.go @@ -18,7 +18,7 @@ const ( nodeAddr5 = "myaddress5" ) -func TestNodeService(t *testing.T) { +func TestWorkerRegistry(t *testing.T) { tx, err := SQLTransactor.BeginTx(context.Background(), true) require.NoError(t, err, "getting transaction") @@ -29,34 +29,34 @@ func TestNodeService(t *testing.T) { } }() - nodeSvc := sqldb.NewNodeService(nil) + workerReg := sqldb.NewWorkerRegistry(nil) - err = nodeSvc.CreateNode(tx, dax.Address(""), &dax.Node{Address: nodeAddr, RoleTypes: []dax.RoleType{"compute"}}) + err = workerReg.AddWorker(tx, &dax.Node{Address: nodeAddr, RoleTypes: []dax.RoleType{dax.RoleTypeCompute}}) require.NoError(t, err) - node, err := nodeSvc.ReadNode(tx, nodeAddr) + node, err := workerReg.Worker(tx, nodeAddr) require.NoError(t, err) require.EqualValues(t, nodeAddr, node.Address) require.EqualValues(t, 1, len(node.RoleTypes)) require.EqualValues(t, "compute", node.RoleTypes[0]) - err = nodeSvc.CreateNode(tx, dax.Address(""), &dax.Node{Address: nodeAddr2, RoleTypes: []dax.RoleType{"translate", "compute"}}) + err = workerReg.AddWorker(tx, &dax.Node{Address: nodeAddr2, RoleTypes: []dax.RoleType{dax.RoleTypeTranslate, dax.RoleTypeCompute}}) require.NoError(t, err, "create node 2") - err = nodeSvc.CreateNode(tx, dax.Address(""), &dax.Node{Address: nodeAddr3, RoleTypes: []dax.RoleType{"compute"}}) + err = workerReg.AddWorker(tx, &dax.Node{Address: nodeAddr3, RoleTypes: []dax.RoleType{dax.RoleTypeCompute}}) require.NoError(t, err, "create node 3") - nodes, err := nodeSvc.Nodes(tx) + nodes, err := workerReg.Workers(tx) require.NoError(t, err) assert.EqualValues(t, 3, len(nodes)) for _, node := range nodes { assert.Contains(t, node.RoleTypes, dax.RoleType("compute"), "node should have compute role but is: %+v", node) } - err = nodeSvc.DeleteNode(tx, nodeAddr2) + err = workerReg.RemoveWorker(tx, nodeAddr2) require.NoError(t, err, "deleting node") - nodes, err = nodeSvc.Nodes(tx) + nodes, err = workerReg.Workers(tx) require.NoError(t, err) require.EqualValues(t, 2, len(nodes)) for _, node := range nodes { diff --git a/dax/controller/balancer/worker_job_test.go b/dax/controller/balancer/worker_job_test.go index f5ada4ab4..8e62d6854 100644 --- a/dax/controller/balancer/worker_job_test.go +++ b/dax/controller/balancer/worker_job_test.go @@ -40,9 +40,14 @@ func TestWorkerJobService(t *testing.T) { wjSvc := sqldb.NewWorkerJobService(nil) qdbid := dax.QualifiedDatabaseID{OrganizationID: orgID, DatabaseID: dbID} + node := &dax.Node{ + Address: nodeAddr, + RoleTypes: []dax.RoleType{role}, + } + // have to create a free worker before you can create a worker job worker - fwSvc := sqldb.NewFreeWorkerService(nil) - err = fwSvc.AddWorkers(tx, role, nodeAddr) + workerReg := sqldb.NewWorkerRegistry(nil) + err = workerReg.AddWorker(tx, node) require.NoError(t, err) err = wjSvc.CreateWorker(tx, role, qdbid, nodeAddr) @@ -65,7 +70,7 @@ func TestWorkerJobService(t *testing.T) { err = fjSvc.CreateJobs(tx, role, qdbid, job1, job2, job3) require.NoError(t, err) - err = wjSvc.CreateJobs(tx, role, qdbid, nodeAddr, job1, job2) + err = wjSvc.AssignWorkerToJobs(tx, role, qdbid, nodeAddr, job1, job2) require.NoError(t, err) jobs, err := wjSvc.ListJobs(tx, role, qdbid, nodeAddr) @@ -85,7 +90,7 @@ func TestWorkerJobService(t *testing.T) { require.NoError(t, err) require.ElementsMatch(t, dax.Addresses{nodeAddr}, addrs) - err = wjSvc.CreateJobs(tx, role, qdbid, nodeAddr, job3) + err = wjSvc.AssignWorkerToJobs(tx, role, qdbid, nodeAddr, job3) require.NoError(t, err) jcs, err := wjSvc.JobCounts(tx, role, qdbid, nodeAddr) @@ -110,7 +115,7 @@ func TestWorkerJobService(t *testing.T) { dk := wjSvc.DatabaseForWorker(tx, nodeAddr) require.EqualValues(t, "db__orgid__blah", dk) - err = wjSvc.DeleteWorker(tx, role, qdbid, nodeAddr) + err = wjSvc.ReleaseWorkers(tx, nodeAddr) require.NoError(t, err) addrs, err = wjSvc.ListWorkers(tx, role, qdbid) diff --git a/dax/controller/controller.go b/dax/controller/controller.go index 5975edcdf..68e3b2872 100644 --- a/dax/controller/controller.go +++ b/dax/controller/controller.go @@ -26,7 +26,7 @@ const ( // Ensure type implements interface. var _ computer.Registrar = (*Controller)(nil) var _ dax.Schemar = (*Controller)(nil) -var _ dax.NodeService = (*Controller)(nil) +var _ dax.WorkerRegistry = (*Controller)(nil) type Controller struct { // Schemar is used by the controller to get table, and other schema, @@ -88,7 +88,7 @@ func New(cfg Config) *Controller { // Poller. pollerCfg := poller.Config{ AddressManager: c, - NodeService: c, + WorkerRegistry: c, NodePoller: poller.NewHTTPNodePoller(logr), PollInterval: cfg.PollInterval, Logger: logr, @@ -671,7 +671,7 @@ func (c *Controller) DropDatabase(ctx context.Context, qdbid dax.QualifiedDataba for worker := range workerSet { addrs = append(addrs, worker) } - if err := c.Balancer.FreeWorkers(tx, addrs...); err != nil { + if err := c.Balancer.ReleaseWorkers(tx, addrs...); err != nil { return errors.Wrap(err, "freeing workers") } @@ -1719,7 +1719,7 @@ func (c *Controller) ComputeNodes(ctx context.Context, qtid dax.QualifiedTableID if err != nil { return nil, errors.Wrap(err, "getting nodes for table") } - computeNodes, err := c.assignedToComputeNodes(assignedNodes) + computeNodes, err := assignedToComputeNodes(assignedNodes) if err != nil { return nil, errors.Wrap(err, "converting assigned to compute nodes") } @@ -1731,13 +1731,13 @@ func (c *Controller) ComputeNodes(ctx context.Context, qtid dax.QualifiedTableID return nil, errors.Wrap(err, "getting compute nodes read or write") } - return c.assignedToComputeNodes(assignedNodes) + return assignedToComputeNodes(assignedNodes) } // assignedToComputeNodes converts the provided []dax.AssignedNode to // []dax.ComputeNode. If any of the assigned nodes are not for RoleType // "compute", an error will be returned. -func (c *Controller) assignedToComputeNodes(nodes []dax.AssignedNode) ([]dax.ComputeNode, error) { +func assignedToComputeNodes(nodes []dax.AssignedNode) ([]dax.ComputeNode, error) { computeNodes := make([]dax.ComputeNode, 0) for _, node := range nodes { @@ -1781,7 +1781,7 @@ func (c *Controller) TranslateNodes(ctx context.Context, qtid dax.QualifiedTable if err != nil { return nil, errors.Wrap(err, "getting nodes for table") } - translateNodes, err := c.assignedToTranslateNodes(assignedNodes) + translateNodes, err := assignedToTranslateNodes(assignedNodes) if err != nil { return nil, errors.Wrap(err, "converting assigned to translate nodes") } @@ -1793,7 +1793,7 @@ func (c *Controller) TranslateNodes(ctx context.Context, qtid dax.QualifiedTable return nil, errors.Wrap(err, "getting translate nodes read or write") } - return c.assignedToTranslateNodes(assignedNodes) + return assignedToTranslateNodes(assignedNodes) } func (c *Controller) nodesForTable(tx dax.Transaction, roleType dax.RoleType, qtid dax.QualifiedTableID) ([]dax.AssignedNode, error) { @@ -1814,7 +1814,7 @@ func (c *Controller) nodesForTable(tx dax.Transaction, roleType dax.RoleType, qt // assignedToTranslateNodes converts the provided []dax.AssignedNode to // []dax.TranslateNode. If any of the assigned nodes are not for RoleType // "translate", an error will be returned. -func (c *Controller) assignedToTranslateNodes(nodes []dax.AssignedNode) ([]dax.TranslateNode, error) { +func assignedToTranslateNodes(nodes []dax.AssignedNode) ([]dax.TranslateNode, error) { translateNodes := make([]dax.TranslateNode, 0) for _, node := range nodes { @@ -1878,7 +1878,7 @@ func (c *Controller) IngestPartition(ctx context.Context, qtid dax.QualifiedTabl } } - translateNodes, err := c.assignedToTranslateNodes(nodes) + translateNodes, err := assignedToTranslateNodes(nodes) if err != nil { return "", errors.Wrap(err, "converting assigned to translate nodes") } @@ -1960,7 +1960,7 @@ func (c *Controller) IngestShard(ctx context.Context, qtid dax.QualifiedTableID, retryAsWrite = true } - computeNodes, err := c.assignedToComputeNodes(nodes) + computeNodes, err := assignedToComputeNodes(nodes) if err != nil { return "", errors.Wrap(err, "converting assigned to compute nodes") } @@ -2207,18 +2207,18 @@ func (c *Controller) TableID(ctx context.Context, qdbid dax.QualifiedDatabaseID, return c.Schemar.TableID(tx, qdbid, name) } -// NodeService +// WorkerRegistry -func (c *Controller) CreateNode(context.Context, dax.Address, *dax.Node) error { - return errors.Errorf("Controller.CreateNode() not implemented") +func (c *Controller) AddWorker(context.Context, dax.Address, *dax.Node) error { + return errors.Errorf("Controller.AddWorker() not implemented") } -func (c *Controller) ReadNode(context.Context, dax.Address) (*dax.Node, error) { - return nil, errors.Errorf("Controller.ReadNode() not implemented") +func (c *Controller) Worker(context.Context, dax.Address) (*dax.Node, error) { + return nil, errors.Errorf("Controller.Worker() not implemented") } -func (c *Controller) DeleteNode(context.Context, dax.Address) error { +func (c *Controller) RemoveWorker(context.Context, dax.Address) error { return errors.Errorf("Controller.DeleteNode() not implemented") } -func (c *Controller) Nodes(ctx context.Context) ([]*dax.Node, error) { +func (c *Controller) Workers(ctx context.Context) ([]*dax.Node, error) { tx, err := c.Transactor.BeginTx(ctx, false) if err != nil { return nil, errors.Wrap(err, "beginning tx") diff --git a/dax/controller/node.go b/dax/controller/node.go deleted file mode 100644 index db64fe438..000000000 --- a/dax/controller/node.go +++ /dev/null @@ -1,36 +0,0 @@ -package controller - -import ( - "github.com/featurebasedb/featurebase/v3/dax" -) - -// NodeService represents a service for managing Nodes. -type NodeService interface { - CreateNode(dax.Transaction, dax.Address, *dax.Node) error - ReadNode(dax.Transaction, dax.Address) (*dax.Node, error) - DeleteNode(dax.Transaction, dax.Address) error - Nodes(dax.Transaction) ([]*dax.Node, error) -} - -// Ensure type implements interface. -var _ NodeService = &nopNodeService{} - -// nopNoder is a no-op implementation of the Noder interface. -type nopNodeService struct{} - -func NewNopNodeService() *nopNodeService { - return &nopNodeService{} -} - -func (n *nopNodeService) CreateNode(dax.Transaction, dax.Address, *dax.Node) error { - return nil -} -func (n *nopNodeService) ReadNode(dax.Transaction, dax.Address) (*dax.Node, error) { - return nil, nil -} -func (n *nopNodeService) DeleteNode(dax.Transaction, dax.Address) error { - return nil -} -func (n *nopNodeService) Nodes(dax.Transaction) ([]*dax.Node, error) { - return []*dax.Node{}, nil -} diff --git a/dax/controller/poller/config.go b/dax/controller/poller/config.go index 3afdc4cd9..a837e7fc8 100644 --- a/dax/controller/poller/config.go +++ b/dax/controller/poller/config.go @@ -9,7 +9,7 @@ import ( type Config struct { AddressManager dax.AddressManager - NodeService dax.NodeService + WorkerRegistry dax.WorkerRegistry NodePoller NodePoller PollInterval time.Duration Logger logger.Logger diff --git a/dax/controller/poller/poller.go b/dax/controller/poller/poller.go index 7725d86cb..8cb83fb39 100644 --- a/dax/controller/poller/poller.go +++ b/dax/controller/poller/poller.go @@ -16,7 +16,7 @@ type Poller struct { addressManager dax.AddressManager - nodeService dax.NodeService + workerRegistry dax.WorkerRegistry nodePoller NodePoller pollInterval time.Duration @@ -30,7 +30,7 @@ type Poller struct { func New(cfg Config) *Poller { p := &Poller{ addressManager: dax.NewNopAddressManager(), - nodeService: dax.NewNopNodeService(), + workerRegistry: dax.NewNopWorkerRegistry(), nodePoller: NewNopNodePoller(), pollInterval: time.Second, logger: logger.NopLogger, @@ -40,8 +40,8 @@ func New(cfg Config) *Poller { if cfg.AddressManager != nil { p.addressManager = cfg.AddressManager } - if cfg.NodeService != nil { - p.nodeService = cfg.NodeService + if cfg.WorkerRegistry != nil { + p.workerRegistry = cfg.WorkerRegistry } if cfg.NodePoller != nil { p.nodePoller = cfg.NodePoller @@ -57,7 +57,7 @@ func New(cfg Config) *Poller { } func (p *Poller) Addresses() []dax.Address { - nodes, err := p.nodeService.Nodes(context.Background()) + nodes, err := p.workerRegistry.Workers(context.Background()) if err != nil { p.logger.Errorf("POLLER: unable to get nodes from node service: %v", err) } diff --git a/dax/controller/poller/poller_test.go b/dax/controller/poller/poller_test.go index 6c3688a74..ca910275d 100644 --- a/dax/controller/poller/poller_test.go +++ b/dax/controller/poller/poller_test.go @@ -25,7 +25,7 @@ import ( func TestPoller(t *testing.T) { ctx := context.Background() - nodeService := newMemNodeService() + workerRegistry := newMemWorkerRegistry() // node 1 node1 := newMockNode(t, "health", 0) @@ -44,7 +44,7 @@ func TestPoller(t *testing.T) { } // manager - manager := newMockManager(t, ctx, "deregister-nodes", nodeService) + manager := newMockManager(t, ctx, "deregister-nodes", workerRegistry) defer manager.Close() managerAddr := dax.Address(manager.URL()) @@ -52,7 +52,7 @@ func TestPoller(t *testing.T) { cfg := poller.Config{ AddressManager: controllerhttp.NewAddressManager(managerAddr), NodePoller: poller.NewHTTPNodePoller(logger.NopLogger), - NodeService: nodeService, + WorkerRegistry: workerRegistry, } p := poller.New(cfg) @@ -62,9 +62,9 @@ func TestPoller(t *testing.T) { close(done) }() - // Add nodes to nodeService so they are available to the poller. - nodeService.CreateNode(ctx, addr1, daxNode1) - nodeService.CreateNode(ctx, addr2, daxNode2) + // Add workers to workerRegistry so they are available to the poller. + workerRegistry.AddWorker(ctx, addr1, daxNode1) + workerRegistry.AddWorker(ctx, addr2, daxNode2) go p.Run() defer p.Stop() @@ -83,13 +83,13 @@ type mockManager struct { t *testing.T server *httptest.Server - nodeService dax.NodeService + workerRegistry dax.WorkerRegistry } -func newMockManager(t *testing.T, ctx context.Context, deregisterPath string, nodeService dax.NodeService) *mockManager { +func newMockManager(t *testing.T, ctx context.Context, deregisterPath string, wr dax.WorkerRegistry) *mockManager { mm := &mockManager{ - t: t, - nodeService: nodeService, + t: t, + workerRegistry: wr, } // deregister is a function used in this mock to remove the address from the @@ -97,7 +97,7 @@ func newMockManager(t *testing.T, ctx context.Context, deregisterPath string, no // the Poller. deregister := func(addrs ...dax.Address) { for _, addr := range addrs { - mm.nodeService.DeleteNode(context.Background(), addr) + mm.workerRegistry.RemoveWorker(context.Background(), addr) } } @@ -176,25 +176,25 @@ func (m *mockNode) Close() { } } -type memNodeService struct { +type memWorkerRegistry struct { mu sync.RWMutex addresses map[dax.Address]*dax.Node } -func newMemNodeService() *memNodeService { - return &memNodeService{ +func newMemWorkerRegistry() *memWorkerRegistry { + return &memWorkerRegistry{ addresses: make(map[dax.Address]*dax.Node), } } -func (m *memNodeService) CreateNode(ctx context.Context, addr dax.Address, node *dax.Node) error { +func (m *memWorkerRegistry) AddWorker(ctx context.Context, addr dax.Address, node *dax.Node) error { m.mu.Lock() defer m.mu.Unlock() m.addresses[addr] = node return nil } -func (m *memNodeService) ReadNode(ctx context.Context, addr dax.Address) (*dax.Node, error) { +func (m *memWorkerRegistry) Worker(ctx context.Context, addr dax.Address) (*dax.Node, error) { m.mu.RLock() defer m.mu.RUnlock() node, ok := m.addresses[addr] @@ -204,14 +204,14 @@ func (m *memNodeService) ReadNode(ctx context.Context, addr dax.Address) (*dax.N return node, nil } -func (m *memNodeService) DeleteNode(ctx context.Context, addr dax.Address) error { +func (m *memWorkerRegistry) RemoveWorker(ctx context.Context, addr dax.Address) error { m.mu.Lock() defer m.mu.Unlock() delete(m.addresses, addr) return nil } -func (m *memNodeService) Nodes(ctx context.Context) ([]*dax.Node, error) { +func (m *memWorkerRegistry) Workers(ctx context.Context) ([]*dax.Node, error) { m.mu.RLock() defer m.mu.RUnlock() diff --git a/dax/controller/sqldb/balancer.go b/dax/controller/sqldb/balancer.go index bd5e1f680..9b7816da8 100644 --- a/dax/controller/sqldb/balancer.go +++ b/dax/controller/sqldb/balancer.go @@ -11,7 +11,7 @@ func NewBalancer(log logger.Logger) *balancer.Balancer { fjs := NewFreeJobService(log) wjs := NewWorkerJobService(log) fws := NewFreeWorkerService(log) - ns := NewNodeService(log) + ns := NewWorkerRegistry(log) return balancer.New(ns, fjs, wjs, fws, schemar, log) } diff --git a/dax/controller/sqldb/freejob.go b/dax/controller/sqldb/freejob.go index d578111e1..505703bc0 100644 --- a/dax/controller/sqldb/freejob.go +++ b/dax/controller/sqldb/freejob.go @@ -102,8 +102,9 @@ func (fj *freeJobService) ListJobs(tx dax.Transaction, roleType dax.RoleType, qd return djs, nil } -// MergeJobs - AFAICT this means "mark these jobs as free" -func (fj *freeJobService) MergeJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, jobs dax.Jobs) error { +// MarkJobsAsFree disassociates any worker that was previously assigned to this +// job. +func (fj *freeJobService) MarkJobsAsFree(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, jobs dax.Jobs) error { dt, ok := tx.(*DaxTransaction) if !ok { return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") diff --git a/dax/controller/sqldb/freeworker.go b/dax/controller/sqldb/freeworker.go index 3ba26ddf3..4ddc0a2e8 100644 --- a/dax/controller/sqldb/freeworker.go +++ b/dax/controller/sqldb/freeworker.go @@ -1,6 +1,8 @@ package sqldb import ( + "fmt" + "github.com/featurebasedb/featurebase/v3/dax" "github.com/featurebasedb/featurebase/v3/dax/controller/balancer" "github.com/featurebasedb/featurebase/v3/dax/models" @@ -21,34 +23,6 @@ type freeWorkerService struct { log logger.Logger } -func (fw *freeWorkerService) AddWorkers(tx dax.Transaction, roleType dax.RoleType, addrs ...dax.Address) error { - dt, ok := tx.(*DaxTransaction) - if !ok { - return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") - } - - workers := make(models.Workers, len(addrs)) - for i, addr := range addrs { - workers[i] = models.Worker{ - Address: addr, - Role: roleType, - } - } - - err := dt.C.Create(workers) - return errors.Wrap(err, "creating workers") -} - -func (fw *freeWorkerService) RemoveWorker(tx dax.Transaction, roleType dax.RoleType, addr dax.Address) error { - dt, ok := tx.(*DaxTransaction) - if !ok { - return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") - } - - err := dt.C.RawQuery("DELETE from workers where database_id is null and role = ? and address = ?", roleType, addr).Exec() - return errors.Wrap(err, "deleting") -} - func (fw *freeWorkerService) PopWorkers(tx dax.Transaction, roleType dax.RoleType, num int) ([]dax.Address, error) { dt, ok := tx.(*DaxTransaction) if !ok { @@ -58,7 +32,8 @@ func (fw *freeWorkerService) PopWorkers(tx dax.Transaction, roleType dax.RoleTyp results := make([]struct { Address dax.Address `db:"address"` }, 0, num) - err := dt.C.RawQuery("select address from workers where role = ? and database_id is NULL limit ?", roleType, num).All(&results) + sel := fmt.Sprintf("select address from workers where role_%s = true and database_id is NULL limit ?", roleType) + err := dt.C.RawQuery(sel, num).All(&results) if err != nil { return nil, errors.Wrap(err, "querying") } @@ -81,9 +56,10 @@ func (fw *freeWorkerService) ListWorkers(tx dax.Transaction, roleType dax.RoleTy } workers := make(models.Workers, 0) - err := dt.C.Select("address").Where("role = ? and database_id is NULL", roleType).Order("address asc").All(&workers) + where := fmt.Sprintf("role_%s = true and database_id is NULL", roleType) + err := dt.C.Select("address").Where(where).Order("address asc").All(&workers) if err != nil { - return nil, errors.Wrap(err, "querying for workers") + return nil, errors.Wrap(err, "querying for free workers") } ret := make(dax.Addresses, len(workers)) diff --git a/dax/controller/sqldb/node.go b/dax/controller/sqldb/node.go deleted file mode 100644 index 9871757a5..000000000 --- a/dax/controller/sqldb/node.go +++ /dev/null @@ -1,116 +0,0 @@ -package sqldb - -import ( - "github.com/featurebasedb/featurebase/v3/dax" - "github.com/featurebasedb/featurebase/v3/dax/controller" - "github.com/featurebasedb/featurebase/v3/dax/models" - "github.com/featurebasedb/featurebase/v3/logger" - "github.com/pkg/errors" -) - -var _ controller.NodeService = (*nodeService)(nil) - -func NewNodeService(log logger.Logger) *nodeService { - if log == nil { - log = logger.NopLogger - } - return &nodeService{ - log: log, - } -} - -type nodeService struct { - log logger.Logger -} - -func (n *nodeService) CreateNode(tx dax.Transaction, addr dax.Address, node *dax.Node) error { - dt, ok := tx.(*DaxTransaction) - if !ok { - return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") - } - - mnode := &models.Node{Address: node.Address} - err := dt.C.Create(mnode) - if err != nil { - return errors.Wrap(err, "creating node") - } - - nodeRoles := make(models.NodeRoles, len(node.RoleTypes)) - mnode.NodeRoles = nodeRoles - for i, rt := range node.RoleTypes { - mnode.NodeRoles[i] = models.NodeRole{ - NodeID: mnode.ID, - Role: rt, - } - } - err = dt.C.Create(&(mnode.NodeRoles)) - if err != nil { - return errors.Wrap(err, "creating node roles") - } - - return nil -} - -func (n *nodeService) ReadNode(tx dax.Transaction, addr dax.Address) (*dax.Node, error) { - dt, ok := tx.(*DaxTransaction) - if !ok { - return nil, dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") - } - - node := &models.Node{} - err := dt.C.Eager().Where("address = ?", addr).First(node) - if err != nil { - return nil, errors.Wrap(err, "getting node") - } - - roleTypes := make([]dax.RoleType, len(node.NodeRoles)) - for i, nr := range node.NodeRoles { - roleTypes[i] = nr.Role - } - - return &dax.Node{ - Address: node.Address, - RoleTypes: roleTypes, - }, nil -} - -func (n *nodeService) DeleteNode(tx dax.Transaction, addr dax.Address) error { - dt, ok := tx.(*DaxTransaction) - if !ok { - return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") - } - - node := &models.Node{} - err := dt.C.Eager().Where("address = ?", addr).First(node) - if isNoRowsError(err) { - return nil - } else if err != nil { - return errors.Wrap(err, "finding node") - } - - err = dt.C.Destroy(node) - return errors.Wrap(err, "destroying node") -} - -func (n *nodeService) Nodes(tx dax.Transaction) ([]*dax.Node, error) { - dt, ok := tx.(*DaxTransaction) - if !ok { - return nil, dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") - } - - nodes := []*models.Node{} - dt.C.Eager().Order("address asc").All(&nodes) - - ret := make([]*dax.Node, len(nodes)) - for i, node := range nodes { - ret[i] = &dax.Node{ - Address: node.Address, - RoleTypes: make([]dax.RoleType, len(node.NodeRoles)), - } - for j, nr := range node.NodeRoles { - ret[i].RoleTypes[j] = nr.Role - } - } - - return ret, nil -} diff --git a/dax/controller/sqldb/worker.go b/dax/controller/sqldb/worker.go new file mode 100644 index 000000000..99bd13bde --- /dev/null +++ b/dax/controller/sqldb/worker.go @@ -0,0 +1,139 @@ +package sqldb + +import ( + "github.com/featurebasedb/featurebase/v3/dax" + "github.com/featurebasedb/featurebase/v3/dax/controller" + "github.com/featurebasedb/featurebase/v3/dax/models" + "github.com/featurebasedb/featurebase/v3/logger" + "github.com/pkg/errors" +) + +var _ controller.WorkerRegistry = (*workerRegistry)(nil) + +func NewWorkerRegistry(log logger.Logger) *workerRegistry { + if log == nil { + log = logger.NopLogger + } + return &workerRegistry{ + log: log, + } +} + +type workerRegistry struct { + log logger.Logger +} + +func (w *workerRegistry) AddWorker(tx dax.Transaction, node *dax.Node) error { + dt, ok := tx.(*DaxTransaction) + if !ok { + return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") + } + + workers := models.Workers{} + + // Determine if a worker for this address already exists. We use `All()` + // here instead of `First()` because `First()` returns an error if there's + // no match. + if err := dt.C.Where("address = ?", node.Address).All(&workers); err != nil { + return errors.Wrapf(err, "getting workers by address: %s", node.Address) + } + + switch len(workers) { + case 0: + // Continue on to create. + case 1: + // Since a worker for this address already exists, just update it and + // return. + worker := workers[0] + for _, roleType := range node.RoleTypes { + if err := worker.SetRole(roleType); err != nil { + return errors.Wrapf(err, "setting role: %s", roleType) + } + } + return dt.C.Update(worker) + default: + return errors.Errorf("found more than one worker for address: %s", node.Address) + } + + worker := &models.Worker{ + Address: node.Address, + } + for _, roleType := range node.RoleTypes { + if err := worker.SetRole(roleType); err != nil { + return errors.Wrapf(err, "setting role: %s", roleType) + } + } + + return dt.C.Create(worker) +} + +func (w *workerRegistry) Worker(tx dax.Transaction, addr dax.Address) (*dax.Node, error) { + dt, ok := tx.(*DaxTransaction) + if !ok { + return nil, dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") + } + + worker := &models.Worker{} + err := dt.C.Eager().Where("address = ?", addr).First(worker) + if err != nil { + return nil, errors.Wrapf(err, "getting worker: %s", addr) + } + + return &dax.Node{ + Address: worker.Address, + RoleTypes: workerRoleTypes(worker), + }, nil +} + +func (w *workerRegistry) RemoveWorker(tx dax.Transaction, addr dax.Address) error { + dt, ok := tx.(*DaxTransaction) + if !ok { + return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") + } + + worker := &models.Worker{} + err := dt.C.Eager().Where("address = ?", addr).First(worker) + if isNoRowsError(err) { + return nil + } else if err != nil { + return errors.Wrapf(err, "finding worker: %s", addr) + } + + err = dt.C.Destroy(worker) + return errors.Wrap(err, "destroying worker") +} + +func (w *workerRegistry) Workers(tx dax.Transaction) ([]*dax.Node, error) { + dt, ok := tx.(*DaxTransaction) + if !ok { + return nil, dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") + } + + workers := []*models.Worker{} + dt.C.Eager().Order("address asc").All(&workers) + + ret := make([]*dax.Node, len(workers)) + for i, worker := range workers { + ret[i] = &dax.Node{ + Address: worker.Address, + RoleTypes: workerRoleTypes(worker), + } + } + + return ret, nil +} + +func workerRoleTypes(worker *models.Worker) []dax.RoleType { + roleTypes := make([]dax.RoleType, 0) + if worker.RoleCompute { + roleTypes = append(roleTypes, dax.RoleTypeCompute) + } + if worker.RoleTranslate { + roleTypes = append(roleTypes, dax.RoleTypeTranslate) + } + if worker.RoleQuery { + roleTypes = append(roleTypes, dax.RoleTypeQuery) + } + + return roleTypes +} diff --git a/dax/controller/sqldb/workerjob.go b/dax/controller/sqldb/workerjob.go index 5ca6466f7..5e4eb15e4 100644 --- a/dax/controller/sqldb/workerjob.go +++ b/dax/controller/sqldb/workerjob.go @@ -24,37 +24,59 @@ type workerJobService struct { log logger.Logger } +// WorkersJobs returns all the workers for the database along with the jobs +// associated to each worker, even if the number of jobs is 0. func (w *workerJobService) WorkersJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID) ([]dax.WorkerInfo, error) { dt, ok := tx.(*DaxTransaction) if !ok { return nil, dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") } + // First, get all workers for the database. workers := models.Workers{} - err := dt.C.Eager().Where("role = ? and database_id = ?", roleType, qdbid.DatabaseID).Order("address asc").All(&workers) + sql := fmt.Sprintf("role_%s = true and database_id = ?", roleType) + err := dt.C.Where(sql, qdbid.DatabaseID).Order("address asc").All(&workers) if err != nil { return nil, errors.Wrap(err, "getting workers") } + // Then, get the jobs for each worker. Ideally, we would do this in a single + // sql query, but it wasn't clear how to do an Eager() LeftJoin() where + // there is a where clause condition on the right side of the join (in this + // case, `jobs.role = ?`). ret := make([]dax.WorkerInfo, len(workers)) for i, worker := range workers { ret[i].Address = worker.Address - ret[i].Jobs = make([]dax.Job, len(worker.Jobs)) - for j, job := range worker.Jobs { - ret[i].Jobs[j] = job.Name + jobs, err := jobsForWorker(dt, &worker, roleType) + if err != nil { + return nil, errors.Wrap(err, "getting jobs for worker") } + ret[i].Jobs = jobs } return ret, nil } +func jobsForWorker(dt *DaxTransaction, worker *models.Worker, roleType dax.RoleType) ([]dax.Job, error) { + jobs := models.Jobs{} + if err := dt.C.Where("worker_id = ? and role = ?", worker.ID, roleType).Order("name asc").All(&jobs); err != nil { + return nil, errors.Wrapf(err, "getting jobs for worker: %s", worker.ID) + } + ret := make([]dax.Job, len(jobs)) + for i := range jobs { + ret[i] = jobs[i].Name + } + return ret, nil +} + func (w *workerJobService) WorkerCount(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID) (int, error) { dt, ok := tx.(*DaxTransaction) if !ok { return 0, dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") } worker := &models.Worker{} - cnt, err := dt.C.Where("role = ? and database_id = ?", roleType, qdbid.DatabaseID).Count(worker) + sql := fmt.Sprintf("role_%s = true and database_id = ?", roleType) + cnt, err := dt.C.Where(sql, qdbid.DatabaseID).Count(worker) return cnt, errors.Wrap(err, "getting count") } @@ -65,7 +87,8 @@ func (w *workerJobService) ListWorkers(tx dax.Transaction, roleType dax.RoleType } workers := models.Workers{} - err := dt.C.Select("address").Where("role = ? and database_id = ?", roleType, qdbid.DatabaseID).Order("address asc").All(&workers) + sql := fmt.Sprintf("role_%s = true and database_id = ?", roleType) + err := dt.C.Select("address").Where(sql, qdbid.DatabaseID).Order("address asc").All(&workers) if err != nil { return nil, errors.Wrap(err, "getting workers") } @@ -85,12 +108,13 @@ func (w *workerJobService) CreateWorker(tx dax.Transaction, roleType dax.RoleTyp } worker := &models.Worker{} - err := dt.C.RawQuery("UPDATE workers SET database_id = ? WHERE role = ? and address = ? RETURNING workers.ID", qdbid.DatabaseID, roleType, addr).First(worker) + sql := fmt.Sprintf("UPDATE workers SET database_id = ? WHERE role_%s = true and address = ? RETURNING workers.ID", roleType) + err := dt.C.RawQuery(sql, qdbid.DatabaseID, addr).First(worker) return errors.Wrap(err, "associating worker to database") } -func (w *workerJobService) FreeWorkers(tx dax.Transaction, addrs ...dax.Address) error { +func (w *workerJobService) ReleaseWorkers(tx dax.Transaction, addrs ...dax.Address) error { dt, ok := tx.(*DaxTransaction) if !ok { return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") @@ -104,34 +128,17 @@ func (w *workerJobService) FreeWorkers(tx dax.Transaction, addrs ...dax.Address) return errors.Wrap(err, "updating workers") } -func (w *workerJobService) DeleteWorker(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address) error { +func (w *workerJobService) AssignWorkerToJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address, job ...dax.Job) error { dt, ok := tx.(*DaxTransaction) if !ok { return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") } worker := &models.Worker{} - err := dt.C.Where("address = ? and role = ? and database_id = ?", addr, roleType, qdbid.DatabaseID).First(worker) - if isNoRowsError(err) { - return nil - } else if err != nil { - return errors.Wrap(err, "getting worker") - } - - err = dt.C.Destroy(worker) - return errors.Wrap(err, "deleting worker") -} - -func (w *workerJobService) CreateJobs(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address, job ...dax.Job) error { - dt, ok := tx.(*DaxTransaction) - if !ok { - return dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") - } - - worker := &models.Worker{} - err := dt.C.Where("address = ? and role = ?", addr, roleType).First(worker) + sql := fmt.Sprintf("address = ? and role_%s = true", roleType) + err := dt.C.Where(sql, addr).First(worker) if err != nil { - return errors.Wrap(err, "getting worker") + return errors.Wrapf(err, "getting worker: (%s) %s", roleType, addr) } jobs := models.Jobs{} @@ -140,35 +147,35 @@ func (w *workerJobService) CreateJobs(tx dax.Transaction, roleType dax.RoleType, return errors.Wrap(err, "updating jobs") } - // create jobs not in "jobs" - toCreate := jobsNotUpdated(job, jobs, worker) + // Assign jobs not in "jobs", and therefore didn't get updated by the + // previous sql statement. + toBeAssigned := jobsNotAssigned(job, jobs, roleType, worker) - err = dt.C.Create(toCreate) - if err != nil { + if err := dt.C.Create(toBeAssigned); err != nil { return errors.Wrap(err, "creating jobs") } - return errors.Wrap(err, "creating jobs") + return nil } -func jobsNotUpdated(incomingJobs []dax.Job, created models.Jobs, worker *models.Worker) (toCreate models.Jobs) { +func jobsNotAssigned(incomingJobs []dax.Job, assigned models.Jobs, roleType dax.RoleType, worker *models.Worker) (toBeAssigned models.Jobs) { outer: for _, incJob := range incomingJobs { - for _, createdJob := range created { - if createdJob.Name == incJob { + for _, assignedJob := range assigned { + if assignedJob.Name == incJob { continue outer } } - toCreate = append(toCreate, + toBeAssigned = append(toBeAssigned, models.Job{ Name: incJob, - Role: worker.Role, + Role: roleType, DatabaseID: dax.DatabaseID(worker.DatabaseID.String), Worker: worker, }, ) } - return toCreate + return toBeAssigned } func (w *workerJobService) DeleteJob(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr dax.Address, job dax.Job) error { @@ -178,7 +185,8 @@ func (w *workerJobService) DeleteJob(tx dax.Transaction, roleType dax.RoleType, } worker := &models.Worker{} - err := dt.C.Select("id").Where("role = ? and database_id = ? and address = ?", roleType, qdbid.DatabaseID, addr).First(worker) + sql := fmt.Sprintf("role_%s = true and database_id = ? and address = ?", roleType) + err := dt.C.Select("id").Where(sql, qdbid.DatabaseID, addr).First(worker) if err != nil { return errors.Wrap(err, "getting worker") } @@ -227,7 +235,7 @@ func (w *workerJobService) DeleteJobsForTable(tx dax.Transaction, roleType dax.R return idiffs, errors.Wrap(err, "deleting jobs") } -func (w *workerJobService) JobCounts(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addr ...dax.Address) (map[dax.Address]int, error) { +func (w *workerJobService) JobCounts(tx dax.Transaction, roleType dax.RoleType, qdbid dax.QualifiedDatabaseID, addrs ...dax.Address) (map[dax.Address]int, error) { dt, ok := tx.(*DaxTransaction) if !ok { return nil, dax.NewErrInvalidTransaction("*sqldb.DaxTransaction") @@ -238,18 +246,23 @@ func (w *workerJobService) JobCounts(tx dax.Transaction, roleType dax.RoleType, Count int `db:"count"` }{} var err error - if len(addr) == 0 { + if len(addrs) == 0 { qstring := `select address, count(*) as count from workers w inner join jobs j on j.worker_id = w.id - where w.database_id = ? and w.role = ? + where w.database_id = ? and w.role_%s = true + and j.role = ? group by w.address` - err = dt.C.RawQuery(qstring, qdbid.DatabaseID, roleType).All(&results) + sql := fmt.Sprintf(qstring, roleType) + err = dt.C.RawQuery(sql, qdbid.DatabaseID, roleType).All(&results) } else { qstring := `select address, count(*) as count from workers w inner join jobs j on j.worker_id = w.id - where w.address in (?) and w.database_id = ? and w.role = ? + where w.database_id = ? and w.role_%s = true + and j.role = ? + and w.address in (?) group by w.address` - err = dt.C.RawQuery(qstring, addr, qdbid.DatabaseID, roleType).All(&results) + sql := fmt.Sprintf(qstring, roleType) + err = dt.C.RawQuery(sql, qdbid.DatabaseID, roleType, addrs).All(&results) } if err != nil { return nil, errors.Wrap(err, "querying for jobs") @@ -269,20 +282,15 @@ func (w *workerJobService) ListJobs(tx dax.Transaction, roleType dax.RoleType, q } worker := &models.Worker{} - err := dt.C.Eager().Where("role = ? and database_id = ? and address = ?", roleType, qdbid.DatabaseID, addr).First(worker) + sql := fmt.Sprintf("role_%s = true and database_id = ? and address = ?", roleType) + err := dt.C.Where(sql, qdbid.DatabaseID, addr).First(worker) if isNoRowsError(err) { return nil, nil } else if err != nil { return nil, errors.Wrap(err, "getting worker") } - ret := make(dax.Jobs, len(worker.Jobs)) - // jobs are ordered by "name asc" defined on the worker model. - for i, job := range worker.Jobs { - ret[i] = job.Name - } - - return ret, nil + return jobsForWorker(dt, worker, roleType) } func (w *workerJobService) DatabaseForWorker(tx dax.Transaction, addr dax.Address) dax.DatabaseKey { diff --git a/dax/controller/sqldb/workerjob_test.go b/dax/controller/sqldb/workerjob_test.go index dd270ccb0..f583d7d4c 100644 --- a/dax/controller/sqldb/workerjob_test.go +++ b/dax/controller/sqldb/workerjob_test.go @@ -21,19 +21,19 @@ func TestJobsNotUpdated(t *testing.T) { }, } - toCreate := jobsNotUpdated(incJobs, created, &models.Worker{ - ID: u2, - Role: "compute", - DatabaseID: nulls.NewString("dbid"), + toCreate := jobsNotAssigned(incJobs, created, dax.RoleTypeCompute, &models.Worker{ + ID: u2, + RoleCompute: true, + DatabaseID: nulls.NewString("dbid"), }) require.Equal(t, 3, len(toCreate)) // Test when 0 jobs are updated - toCreate = jobsNotUpdated(incJobs, models.Jobs{}, &models.Worker{ - ID: u2, - Role: "compute", - DatabaseID: nulls.NewString("dbid"), + toCreate = jobsNotAssigned(incJobs, models.Jobs{}, dax.RoleTypeCompute, &models.Worker{ + ID: u2, + RoleCompute: true, + DatabaseID: nulls.NewString("dbid"), }) require.Equal(t, 4, len(toCreate)) diff --git a/dax/controller/worker.go b/dax/controller/worker.go new file mode 100644 index 000000000..a2d2ac8c0 --- /dev/null +++ b/dax/controller/worker.go @@ -0,0 +1,40 @@ +package controller + +import ( + "github.com/featurebasedb/featurebase/v3/dax" +) + +// WorkerRegistry represents a service for managing Nodes. Note that this +// interface mirrors the dax.WorkerRegistry interface, but its methods take +// dax.Transactions rather than Contexts. That's because the dax version of this +// interface is meant to be a the API boundary, where this is an interface for +// use within the Controller. +type WorkerRegistry interface { + AddWorker(dax.Transaction, *dax.Node) error + Worker(dax.Transaction, dax.Address) (*dax.Node, error) + RemoveWorker(dax.Transaction, dax.Address) error + Workers(dax.Transaction) ([]*dax.Node, error) +} + +// Ensure type implements interface. +var _ WorkerRegistry = &nopWorkerRegistry{} + +// nopWorkerRegistry is a no-op implementation of the WorkerRegistry interface. +type nopWorkerRegistry struct{} + +func NewNopWorkerRegistry() *nopWorkerRegistry { + return &nopWorkerRegistry{} +} + +func (n *nopWorkerRegistry) AddWorker(dax.Transaction, *dax.Node) error { + return nil +} +func (n *nopWorkerRegistry) Worker(dax.Transaction, dax.Address) (*dax.Node, error) { + return nil, nil +} +func (n *nopWorkerRegistry) RemoveWorker(dax.Transaction, dax.Address) error { + return nil +} +func (n *nopWorkerRegistry) Workers(dax.Transaction) ([]*dax.Node, error) { + return []*dax.Node{}, nil +} diff --git a/dax/migrations/003_node_to_worker.down.fizz b/dax/migrations/003_node_to_worker.down.fizz new file mode 100644 index 000000000..e69de29bb diff --git a/dax/migrations/003_node_to_worker.up.fizz b/dax/migrations/003_node_to_worker.up.fizz new file mode 100644 index 000000000..ee3d636c9 --- /dev/null +++ b/dax/migrations/003_node_to_worker.up.fizz @@ -0,0 +1,16 @@ +add_column("workers", "role_compute", "bool", {"default": false}) +add_column("workers", "role_translate", "bool", {"default": false}) +add_column("workers", "role_query", "bool", {"default": false}) + +sql("update jobs set worker_id = wc.id from workers wc inner join workers wt on wc.address = wt.address and wc.database_id = wt.database_id and wc.role = 'compute' and wt.role = 'translate' where worker_id = wt.id;") + +sql("update workers set role_compute = true where role = 'compute';") + +sql("update workers set role_translate = true from workers wt where workers.address = wt.address and workers.role = 'compute' and wt.role = 'translate';") + +sql("delete from workers where role = 'translate';") + +drop_column("workers", "role") + +drop_table("node_roles") +drop_table("nodes") \ No newline at end of file diff --git a/dax/models/node.go b/dax/models/node.go deleted file mode 100644 index c347e75ba..000000000 --- a/dax/models/node.go +++ /dev/null @@ -1,56 +0,0 @@ -package models - -import ( - "encoding/json" - "time" - - "github.com/featurebasedb/featurebase/v3/dax" - "github.com/gobuffalo/pop/v6" - "github.com/gobuffalo/validate/v3" - "github.com/gobuffalo/validate/v3/validators" - "github.com/gofrs/uuid" -) - -// Node represents a host or server that is available to work on jobs. -type Node struct { - ID uuid.UUID `json:"id" db:"id"` - Address dax.Address `json:"address" db:"address"` - NodeRoles NodeRoles `json:"node_roles" has_many:"node_roles" order_by:"created_at asc"` - CreatedAt time.Time `json:"created_at" db:"created_at"` - UpdatedAt time.Time `json:"updated_at" db:"updated_at"` -} - -// String is not required by pop and may be deleted -func (t *Node) String() string { - jt, _ := json.MarshalIndent(t, " ", " ") //nolint:errchkjson - return string(jt) -} - -// Nodes is not required by pop and may be deleted -type Nodes []*Node - -// String is not required by pop and may be deleted -func (t Nodes) String() string { - jt, _ := json.MarshalIndent(t, " ", " ") //nolint:errchkjson - return string(jt) -} - -// Validate gets run every time you call a "pop.Validate*" (pop.ValidateAndSave, pop.ValidateAndCreate, pop.ValidateAndUpdate) method. -// This method is not required and may be deleted. -func (t *Node) Validate(tx *pop.Connection) (*validate.Errors, error) { - return validate.Validate( - &validators.StringIsPresent{Field: string(t.Address), Name: "Address"}, - ), nil -} - -// ValidateCreate gets run every time you call "pop.ValidateAndCreate" method. -// This method is not required and may be deleted. -func (t *Node) ValidateCreate(tx *pop.Connection) (*validate.Errors, error) { - return validate.NewErrors(), nil -} - -// ValidateUpdate gets run every time you call "pop.ValidateAndUpdate" method. -// This method is not required and may be deleted. -func (t *Node) ValidateUpdate(tx *pop.Connection) (*validate.Errors, error) { - return validate.NewErrors(), nil -} diff --git a/dax/models/node_roles.go b/dax/models/node_roles.go deleted file mode 100644 index 558e2940e..000000000 --- a/dax/models/node_roles.go +++ /dev/null @@ -1,31 +0,0 @@ -package models - -import ( - "encoding/json" - "time" - - "github.com/featurebasedb/featurebase/v3/dax" - "github.com/gofrs/uuid" -) - -// NodeRole holds information about what types of jobs (roles) each node can perform. -type NodeRole struct { - ID uuid.UUID `json:"id" db:"id"` - NodeID uuid.UUID `json:"node_id" db:"node_id"` - Role dax.RoleType `json:"role" db:"role"` - CreatedAt time.Time `json:"created_at" db:"created_at"` - UpdatedAt time.Time `json:"updated_at" db:"updated_at"` -} - -// String is not required by pop and may be deleted -func (t *NodeRole) String() string { - jt, _ := json.MarshalIndent(t, " ", " ") //nolint:errchkjson - return string(jt) -} - -type NodeRoles []NodeRole - -func (t NodeRoles) String() string { - jt, _ := json.MarshalIndent(t, " ", " ") //nolint:errchkjson - return string(jt) -} diff --git a/dax/models/worker.go b/dax/models/worker.go index c2ddd25f7..dd4b2f249 100644 --- a/dax/models/worker.go +++ b/dax/models/worker.go @@ -5,6 +5,7 @@ import ( "time" "github.com/featurebasedb/featurebase/v3/dax" + "github.com/featurebasedb/featurebase/v3/errors" "github.com/gobuffalo/nulls" "github.com/gobuffalo/pop/v6" "github.com/gobuffalo/validate/v3" @@ -15,13 +16,15 @@ import ( // Worker is a node plus a role that gets assigned to a database and // can be assigned jobs for that database. type Worker struct { - ID uuid.UUID `json:"id" db:"id"` - Address dax.Address `json:"address" db:"address"` - Role dax.RoleType `json:"role" db:"role"` - DatabaseID nulls.String `json:"database_id" db:"database_id"` // this can be empty which means the worker is unassigned - CreatedAt time.Time `json:"created_at" db:"created_at"` - UpdatedAt time.Time `json:"updated_at" db:"updated_at"` - Jobs Jobs `json:"jobs" has_many:"jobs" order_by:"name asc"` + ID uuid.UUID `json:"id" db:"id"` + Address dax.Address `json:"address" db:"address"` + DatabaseID nulls.String `json:"database_id" db:"database_id"` // this can be empty which means the worker is unassigned + CreatedAt time.Time `json:"created_at" db:"created_at"` + UpdatedAt time.Time `json:"updated_at" db:"updated_at"` + Jobs Jobs `json:"jobs" has_many:"jobs" order_by:"name asc"` + RoleCompute bool `json:"role_compute" db:"role_compute"` + RoleTranslate bool `json:"role_translate" db:"role_translate"` + RoleQuery bool `json:"role_query" db:"role_query"` } // String is not required by pop and may be deleted @@ -30,6 +33,23 @@ func (t *Worker) String() string { return string(jt) } +// SetRole applies a dax.RoleType to one of the boolean fields on the Worker +// model. It returns an error if the model does not support that role type. +func (t *Worker) SetRole(role dax.RoleType) error { + switch role { + case dax.RoleTypeCompute: + t.RoleCompute = true + case dax.RoleTypeTranslate: + t.RoleTranslate = true + case dax.RoleTypeQuery: + t.RoleQuery = true + default: + errors.Errorf("invalid role type for worker: %s", role) + } + + return nil +} + // Workers is not required by pop and may be deleted type Workers []Worker diff --git a/dax/role.go b/dax/role.go index 165356611..e38dce7f9 100644 --- a/dax/role.go +++ b/dax/role.go @@ -6,6 +6,11 @@ type RoleType string const ( RoleTypeCompute RoleType = "compute" RoleTypeTranslate RoleType = "translate" + RoleTypeQuery RoleType = "query" +) + +var ( + AllRoleTypes = []RoleType{RoleTypeCompute, RoleTypeTranslate, RoleTypeQuery} ) // RoleTypes is a list of RoleType, used primarily to introduce helper methods diff --git a/dax/test/dax/dax_test.go b/dax/test/dax/dax_test.go index cfa1da94c..df1e56966 100644 --- a/dax/test/dax/dax_test.go +++ b/dax/test/dax/dax_test.go @@ -157,7 +157,8 @@ func TestDAXIntegration(t *testing.T) { "testinsert/test-5", // error messages differ "percentile_test/test-6", // related to TODO in orchestrator.executePercentile "alterTable/alterTableBadTable", // looks like table does not exist is a different error in DAX - "top-tests/test-1", // don't know why this is failing at all + "top-limit-tests/test-2", // don't know why this is failing at all + "top-limit-tests/test-3", // don't know why this is failing at all "delete_tests", "viewtests/drop-view", // drop view does a delete "viewtests/drop-view-if-exists-after-drop", @@ -287,7 +288,7 @@ func TestDAXIntegration(t *testing.T) { // TODO: implement this without a sleep. time.Sleep(5 * time.Second) - // ensure paritions are still covered + // ensure partitions are still covered nodes, err = controllerClient.TranslateNodes(context.Background(), qtid, append(partitions0, partitions1...)...) assert.NoError(t, err) if assert.Len(t, nodes, 1) { @@ -477,7 +478,7 @@ func TestDAXIntegration(t *testing.T) { assert.NoError(t, svcmgr.ControllerStart()) assert.True(t, mc.Healthy(controllerKey)) - // ensure paritions are still covered + // ensure partitions are still covered nodes, err = controllerClient.TranslateNodes(context.Background(), qtid, partitions...) assert.NoError(t, err) if assert.Len(t, nodes, 1) { diff --git a/dax/node.go b/dax/worker.go similarity index 85% rename from dax/node.go rename to dax/worker.go index 62060af7b..e4ad397f4 100644 --- a/dax/node.go +++ b/dax/worker.go @@ -53,34 +53,34 @@ type AssignedNode struct { Role Role `json:"role"` } -// NodeService represents a service for managing Nodes. -type NodeService interface { - CreateNode(context.Context, Address, *Node) error - ReadNode(context.Context, Address) (*Node, error) - DeleteNode(context.Context, Address) error - Nodes(context.Context) ([]*Node, error) +// WorkerRegistry represents a service for managing Workers. +type WorkerRegistry interface { + AddWorker(context.Context, Address, *Node) error + Worker(context.Context, Address) (*Node, error) + RemoveWorker(context.Context, Address) error + Workers(context.Context) ([]*Node, error) } // Ensure type implements interface. -var _ NodeService = &nopNodeService{} +var _ WorkerRegistry = &nopWorkerRegistry{} -// nopNoder is a no-op implementation of the Noder interface. -type nopNodeService struct{} +// nopWorkerRegistry is a no-op implementation of the WorkerRegistry interface. +type nopWorkerRegistry struct{} -func NewNopNodeService() *nopNodeService { - return &nopNodeService{} +func NewNopWorkerRegistry() *nopWorkerRegistry { + return &nopWorkerRegistry{} } -func (n *nopNodeService) CreateNode(context.Context, Address, *Node) error { +func (n *nopWorkerRegistry) AddWorker(context.Context, Address, *Node) error { return nil } -func (n *nopNodeService) ReadNode(context.Context, Address) (*Node, error) { +func (n *nopWorkerRegistry) Worker(context.Context, Address) (*Node, error) { return nil, nil } -func (n *nopNodeService) DeleteNode(context.Context, Address) error { +func (n *nopWorkerRegistry) RemoveWorker(context.Context, Address) error { return nil } -func (n *nopNodeService) Nodes(context.Context) ([]*Node, error) { +func (n *nopWorkerRegistry) Workers(context.Context) ([]*Node, error) { return []*Node{}, nil } diff --git a/go.mod b/go.mod index c218b2e7d..0006daf48 100644 --- a/go.mod +++ b/go.mod @@ -137,6 +137,7 @@ require ( github.com/sourcegraph/annotate v0.0.0-20160123013949-f4cad6c6324d // indirect github.com/sourcegraph/syntaxhighlight v0.0.0-20170531221838-bd320f5d308e // indirect github.com/tinylib/msgp v1.1.2 // indirect + gonum.org/v1/gonum v0.11.0 // indirect ) require ( @@ -197,6 +198,7 @@ require ( github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect github.com/prometheus/common v0.37.0 // indirect github.com/prometheus/procfs v0.8.0 // indirect + github.com/sajari/regression v1.0.1 github.com/sirupsen/logrus v1.9.0 // indirect github.com/soheilhy/cmux v0.1.5 // indirect github.com/spf13/afero v1.6.0 // indirect diff --git a/go.sum b/go.sum index e90e72013..e442c2b33 100644 --- a/go.sum +++ b/go.sum @@ -1056,6 +1056,8 @@ github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQD github.com/ryanuber/columnize v0.0.0-20160712163229-9b3edd62028f/go.mod h1:sm1tb6uqfes/u+d4ooFouqFdy9/2g9QGwK3SQygK0Ts= github.com/ryanuber/columnize v2.1.0+incompatible/go.mod h1:sm1tb6uqfes/u+d4ooFouqFdy9/2g9QGwK3SQygK0Ts= github.com/ryanuber/go-glob v1.0.0/go.mod h1:807d1WSdnB0XRJzKNil9Om6lcp/3a0v4qIHxIXzX/Yc= +github.com/sajari/regression v1.0.1 h1:iTVc6ZACGCkoXC+8NdqH5tIreslDTT/bXxT6OmHR5PE= +github.com/sajari/regression v1.0.1/go.mod h1:NeG/XTW1lYfGY7YV/Z0nYDV/RGh3wxwd1yW46835flM= github.com/samuel/go-zookeeper v0.0.0-20190923202752-2cc03de413da/go.mod h1:gi+0XIa01GRL2eRQVjQkKGqKF3SF9vZR/HnPullcV2E= github.com/santhosh-tekuri/jsonschema/v5 v5.0.0/go.mod h1:FKdcjfQW6rpZSnxxUvEA5H/cDPdvJ/SZJQLWWXWGrZ0= github.com/satori/go.uuid v1.2.0/go.mod h1:dA0hQrYB0VpLJoorglMZABFdXlWrHn1NEOzdhQKdks0= @@ -1684,6 +1686,7 @@ golang.org/x/xerrors v0.0.0-20220609144429-65e65417b02f/go.mod h1:K8+ghG5WaK9qNq gonum.org/v1/gonum v0.0.0-20180816165407-929014505bf4/go.mod h1:Y+Yx5eoAFn32cQvJDxZx5Dpnq+c3wtXuadVZAcxbbBo= gonum.org/v1/gonum v0.8.2/go.mod h1:oe/vMfY3deqTw+1EZJhuvEW2iwGF1bW9wwu7XCu0+v0= gonum.org/v1/gonum v0.11.0 h1:f1IJhK4Km5tBJmaiJXtk/PkL4cdVX6J+tGiM187uT5E= +gonum.org/v1/gonum v0.11.0/go.mod h1:fSG4YDCxxUZQJ7rKsQrj0gMOg00Il0Z96/qMA4bVQhA= gonum.org/v1/netlib v0.0.0-20190313105609-8cb42192e0e0/go.mod h1:wa6Ws7BG/ESfp6dHfk7C6KdzKA7wR7u/rKwOGE66zvw= gonum.org/v1/plot v0.0.0-20190515093506-e2840ee46a6b/go.mod h1:Wt8AAjI+ypCyYX3nZBvf6cAIx93T+c/OS2HFAYskSZc= google.golang.org/api v0.3.1/go.mod h1:6wY9I6uQWHQ8EM57III9mq/AjF+i8G65rmVagqKMtkk= diff --git a/sql3/errors.go b/sql3/errors.go index 6ed3ebcaa..3bc42cf6a 100644 --- a/sql3/errors.go +++ b/sql3/errors.go @@ -13,10 +13,13 @@ const ( ErrUnsupported errors.Code = "ErrUnsupported" ErrCacheKeyNotFound errors.Code = "ErrCacheKeyNotFound" - ErrDuplicateColumn errors.Code = "ErrDuplicateColumn" - ErrUnknownType errors.Code = "ErrUnknownType" - ErrUnknownIdentifier errors.Code = "ErrUnknownIdentifier" + // syntax/semantic errors + ErrDuplicateColumn errors.Code = "ErrDuplicateColumn" + ErrUnknownType errors.Code = "ErrUnknownType" + ErrUnknownIdentifier errors.Code = "ErrUnknownIdentifier" + ErrTopLimitCannotCoexist errors.Code = "ErrTopLimitCannotCoexist" + // type related errors ErrTypeIncompatibleWithBitwiseOperator errors.Code = "ErrTypeIncompatibleWithBitwiseOperator" ErrTypeIncompatibleWithLogicalOperator errors.Code = "ErrTypeIncompatibleWithLogicalOperator" ErrTypeIncompatibleWithEqualityOperator errors.Code = "ErrTypeIncompatibleWithEqualityOperator" @@ -40,8 +43,6 @@ const ( ErrTimeQuantumExpressionExpected errors.Code = "ErrTimeQuantumExpressionExpected" ErrSingleRowExpected errors.Code = "ErrSingleRowExpected" - // type related errors - // decimal ErrDecimalScaleExpected errors.Code = "ErrDecimalScaleExpected" @@ -94,6 +95,9 @@ const ( ErrViewExists errors.Code = "ErrViewExists" ErrViewNotFound errors.Code = "ErrViewNotFound" + ErrModelExists errors.Code = "ErrModelExists" + ErrModelNotFound errors.Code = "ErrModelNotFound" + ErrBadColumnConstraint errors.Code = "ErrBadColumnConstraint" ErrConflictingColumnConstraint errors.Code = "ErrConflictingColumnConstraint" @@ -141,6 +145,9 @@ const ( ErrInvalidDatetimePart errors.Code = "ErrInvalidDatetimePart" ErrOutputValueOutOfRange errors.Code = "ErrOutputValueOutOfRange" ErrDivideByZero errors.Code = "ErrDivideByZero" + + // remote execution + ErrRemoteUnauthorized errors.Code = "ErrRemoteUnauthorized" ) func NewErrDuplicateColumn(line int, col int, column string) error { @@ -164,6 +171,13 @@ func NewErrUnknownIdentifier(line int, col int, ident string) error { ) } +func NewErrErrTopLimitCannotCoexist(line int, col int) error { + return errors.New( + ErrTopLimitCannotCoexist, + fmt.Sprintf("[%d:%d] TOP and LIMIT cannot cannot be used at the same time (TOP will be deprecated in a future release)", line, col), + ) +} + func NewErrInternal(msg string) error { preamble := "internal error" _, filename, line, ok := runtime.Caller(1) @@ -647,6 +661,20 @@ func NewErrViewExists(line, col int, viewName string) error { ) } +func NewErrModelNotFound(line, col int, viewName string) error { + return errors.New( + ErrModelNotFound, + fmt.Sprintf("[%d:%d] model '%s' not found", line, col, viewName), + ) +} + +func NewErrModelExists(line, col int, viewName string) error { + return errors.New( + ErrModelExists, + fmt.Sprintf("[%d:%d] model '%s' already exists", line, col, viewName), + ) +} + func NewErrBadColumnConstraint(line, col int, constraint, columnType string) error { return errors.New( ErrBadColumnConstraint, @@ -876,3 +904,10 @@ func NewErrDivideByZero(line, col int) error { fmt.Sprintf("[%d:%d] divisor is equal to zero", line, col), ) } + +func NewErrRemoteUnauthorized(line, col int, remoteUrl string) error { + return errors.New( + ErrRemoteUnauthorized, + fmt.Sprintf("unauthorized on remote server '%s'", remoteUrl), + ) +} diff --git a/sql3/parser/ast.go b/sql3/parser/ast.go index dd5501105..3c34d0ca3 100644 --- a/sql3/parser/ast.go +++ b/sql3/parser/ast.go @@ -34,11 +34,13 @@ func (*CaseExpr) node() {} func (*CastExpr) node() {} func (*CheckConstraint) node() {} func (*ColumnDefinition) node() {} +func (*CopyStatement) node() {} func (*CommitStatement) node() {} func (*CreateDatabaseStatement) node() {} func (*CreateIndexStatement) node() {} func (*CreateTableStatement) node() {} func (*CreateFunctionStatement) node() {} +func (*CreateModelStatement) node() {} func (*CreateViewStatement) node() {} func (*DateLit) node() {} func (*DefaultConstraint) node() {} @@ -48,6 +50,7 @@ func (*DropIndexStatement) node() {} func (*DropTableStatement) node() {} func (*DropFunctionStatement) node() {} func (*DropViewStatement) node() {} +func (*DropModelStatement) node() {} func (*Exists) node() {} func (*ExplainStatement) node() {} func (*ExprList) node() {} @@ -73,12 +76,14 @@ func (*OnConstraint) node() {} func (*OrderingTerm) node() {} func (*OverClause) node() {} func (*ParenExpr) node() {} +func (*PredictStatement) node() {} func (*SetLiteralExpr) node() {} func (*ParenSource) node() {} func (*PrimaryKeyConstraint) node() {} func (*QualifiedRef) node() {} func (*QualifiedTableName) node() {} func (*Range) node() {} +func (*ReturnStatement) node() {} func (*ReleaseStatement) node() {} func (*ResultColumn) node() {} func (*RollbackStatement) node() {} @@ -111,6 +116,7 @@ func (*AlterTableStatement) stmt() {} func (*AlterViewStatement) stmt() {} func (*AnalyzeStatement) stmt() {} func (*BeginStatement) stmt() {} +func (*CopyStatement) stmt() {} func (*BulkInsertStatement) stmt() {} func (*ShowDatabasesStatement) stmt() {} func (*ShowTablesStatement) stmt() {} @@ -121,6 +127,7 @@ func (*CreateDatabaseStatement) stmt() {} func (*CreateIndexStatement) stmt() {} func (*CreateTableStatement) stmt() {} func (*CreateFunctionStatement) stmt() {} +func (*CreateModelStatement) stmt() {} func (*CreateViewStatement) stmt() {} func (*DeleteStatement) stmt() {} func (*DropDatabaseStatement) stmt() {} @@ -128,9 +135,12 @@ func (*DropIndexStatement) stmt() {} func (*DropTableStatement) stmt() {} func (*DropFunctionStatement) stmt() {} func (*DropViewStatement) stmt() {} +func (*DropModelStatement) stmt() {} +func (*PredictStatement) stmt() {} func (*ExplainStatement) stmt() {} func (*InsertStatement) stmt() {} func (*ReleaseStatement) stmt() {} +func (*ReturnStatement) stmt() {} func (*RollbackStatement) stmt() {} func (*SavepointStatement) stmt() {} func (*SelectStatement) stmt() {} @@ -177,6 +187,8 @@ func CloneStatement(stmt Statement) Statement { return stmt.Clone() case *DropViewStatement: return stmt.Clone() + case *DropModelStatement: + return stmt.Clone() case *ExplainStatement: return stmt.Clone() case *InsertStatement: @@ -2910,6 +2922,35 @@ func (s *DropViewStatement) String() string { return buf.String() } +type DropModelStatement struct { + Drop Pos // position of DROP keyword + Model Pos // position of MODEL keyword + If Pos // position of IF keyword + IfExists Pos // position of EXISTS keyword after IF + Name *Ident // view name +} + +// Clone returns a deep copy of s. +func (s *DropModelStatement) Clone() *DropModelStatement { + if s == nil { + return nil + } + other := *s + other.Name = s.Name.Clone() + return &other +} + +// String returns the string representation of the statement. +func (s *DropModelStatement) String() string { + var buf bytes.Buffer + buf.WriteString("DROP MODEL") + if s.IfExists.IsValid() { + buf.WriteString(" IF EXISTS") + } + fmt.Fprintf(&buf, " %s", s.Name.String()) + return buf.String() +} + type CreateIndexStatement struct { Create Pos // position of CREATE keyword Unique Pos // position of optional UNIQUE keyword @@ -2998,6 +3039,11 @@ func (s *DropIndexStatement) String() string { return buf.String() } +type FunctionOptionDefinition struct { + Name *Ident // option name + OptionExpr Expr // option expression +} + type ParameterDefinition struct { Name *Variable // parameter name Type *Type // data type @@ -3015,8 +3061,11 @@ type CreateFunctionStatement struct { Parameters []*ParameterDefinition // parameters Rparen Pos // position of parameter RParen - Returns Pos // position of RETURNS keyword - ReturnDef *ParameterDefinition // return def + Returns Pos // position of RETURNS keyword + ReturnType *Type // return def + + With Pos // position of WITH keyword + Options []*FunctionOptionDefinition // options As Pos // position of AS keyword @@ -3057,7 +3106,21 @@ func (s *CreateFunctionStatement) String() string { } buf.WriteString(" RETURNS ") - fmt.Fprintf(&buf, "%s %s", s.ReturnDef.Name, s.ReturnDef.Type.Name) + fmt.Fprintf(&buf, "%s", s.ReturnType.String()) + + if s.With.IsValid() { + buf.WriteString(" WITH ") + if len(s.Options) > 0 { + buf.WriteString(" (") + for idx, p := range s.Options { + if idx > 0 { + buf.WriteString(", ") + } + fmt.Fprintf(&buf, "%s %s", p.Name.Name, p.OptionExpr.String()) + } + buf.WriteString(")") + } + } buf.WriteString(" AS BEGIN") for i := range s.Body { @@ -3096,6 +3159,86 @@ func (s *DropFunctionStatement) String() string { return buf.String() } +type ReturnStatement struct { + Return Pos // position of RETURN keyword + ReturnExpr Expr // what we are returning +} + +// Clone returns a deep copy of s. +func (s *ReturnStatement) Clone() *ReturnStatement { + if s == nil { + return nil + } + other := *s + other.ReturnExpr = CloneExpr(s.ReturnExpr) + return &other +} + +func (s *ReturnStatement) String() string { + var buf bytes.Buffer + buf.WriteString("RETURN") + fmt.Fprintf(&buf, " %s", s.ReturnExpr.String()) + return buf.String() +} + +type ModelOptionDefinition struct { + Name *Ident // option name + OptionExpr Expr // option expression +} + +type CreateModelStatement struct { + Create Pos // position of CREATE keyword + Model Pos // position of MODEL keyword + If Pos // position of IF keyword + IfNot Pos // position of NOT keyword after IF + IfNotExists Pos // position of EXISTS keyword after IF NOT + Name *Ident // model name + With Pos // position of WITH keyword + + Options []*ModelOptionDefinition // options + + As Pos // position of AS keyword + + ModelQuery *SelectStatement // model query +} + +// Clone returns a deep copy of s. +func (s *CreateModelStatement) Clone() *CreateModelStatement { + if s == nil { + return nil + } + other := *s + other.Name = s.Name.Clone() + other.ModelQuery = s.ModelQuery.Clone() + return &other +} + +// String returns the string representation of the statement. +func (s *CreateModelStatement) String() string { + var buf bytes.Buffer + buf.WriteString("CREATE MODEL") + if s.IfNotExists.IsValid() { + buf.WriteString(" IF NOT EXISTS") + } + fmt.Fprintf(&buf, " %s", s.Name.String()) + + if len(s.Options) > 0 { + buf.WriteString(" (") + for idx, p := range s.Options { + if idx > 0 { + buf.WriteString(", ") + } + fmt.Fprintf(&buf, "%s %s", p.Name.Name, p.OptionExpr.String()) + } + buf.WriteString(")") + } + + buf.WriteString(" AS ") + buf.WriteString(s.ModelQuery.String()) + + return buf.String() +} + type BulkInsertMapDefinition struct { Name *Ident // map name Type *Type // data type @@ -3645,6 +3788,77 @@ func (c *IndexedColumn) String() string { return c.X.String() } +type CopyStatement struct { + Copy Pos // position of COPY keyword + Source Source // source table + To Pos // position of TO keyword + TargetName *Ident // target table name + Where Pos // position of WHERE keyword + WhereExpr Expr // where clause expression + With Pos // position of WITH keyword + + Url Expr // url for target server + ApiKey Expr // apikey for target server +} + +func (c *CopyStatement) Clone() *CopyStatement { + if c == nil { + return nil + } + other := *c + other.Source = CloneSource(c.Source) + other.TargetName = c.TargetName.Clone() + other.WhereExpr = CloneExpr(c.WhereExpr) + other.Url = CloneExpr(c.Url) + other.ApiKey = CloneExpr(c.ApiKey) + return &other +} + +func (c *CopyStatement) String() string { + var buf bytes.Buffer + + fmt.Fprintf(&buf, "COPY %s to %s", c.Source.String(), c.TargetName.String()) + if c.WhereExpr != nil { + fmt.Fprintf(&buf, " WHERE %s", c.WhereExpr.String()) + } + if c.With.IsValid() { + fmt.Fprintf(&buf, " WITH") + if c.Url != nil { + fmt.Fprintf(&buf, " URL %s", c.Url.String()) + } + if c.ApiKey != nil { + fmt.Fprintf(&buf, " APIKEY %s", c.Url.String()) + } + } + return buf.String() +} + +type PredictStatement struct { + Predict Pos // position of PREDICT keyword + Using Pos // position of USING keyword + ModelName *Ident // model name + + InputQuery *SelectStatement // input query +} + +func (c *PredictStatement) Clone() *PredictStatement { + if c == nil { + return nil + } + other := *c + other.ModelName = c.ModelName.Clone() + other.InputQuery = c.InputQuery.Clone() + return &other +} + +func (c *PredictStatement) String() string { + var buf bytes.Buffer + + fmt.Fprintf(&buf, "PREDICT USING %s", c.ModelName.String()) + fmt.Fprintf(&buf, " %s", c.InputQuery.String()) + return buf.String() +} + type SelectStatement struct { WithClause *WithClause // clause containing CTEs @@ -3685,6 +3899,8 @@ type SelectStatement struct { OrderBy Pos // position of BY keyword after ORDER OrderingTerms []*OrderingTerm // terms of ORDER BY clause + Limit Pos // position of LIMIT keyword + LimitExpr Expr // LIMIT expr } // Clone returns a deep copy of s. @@ -3694,7 +3910,6 @@ func (s *SelectStatement) Clone() *SelectStatement { } other := *s other.WithClause = s.WithClause.Clone() - //other.ValueLists = cloneExprLists(s.ValueLists) other.TopExpr = CloneExpr(s.TopExpr) other.Columns = cloneResultColumns(s.Columns) other.Source = CloneSource(s.Source) @@ -3704,6 +3919,7 @@ func (s *SelectStatement) Clone() *SelectStatement { other.Windows = cloneWindows(s.Windows) other.Compound = s.Compound.Clone() other.OrderingTerms = cloneOrderingTerms(s.OrderingTerms) + other.LimitExpr = CloneExpr(s.LimitExpr) return &other } @@ -3742,29 +3958,10 @@ func (s *SelectStatement) String() string { buf.WriteString(" ") } - /*if len(s.ValueLists) > 0 { - buf.WriteString("VALUES ") - for i, exprs := range s.ValueLists { - if i != 0 { - buf.WriteString(", ") - } - - buf.WriteString("(") - for j, expr := range exprs.Exprs { - if j != 0 { - buf.WriteString(", ") - } - buf.WriteString(expr.String()) - } - buf.WriteString(")") - } - } else {*/ buf.WriteString("SELECT ") if s.Distinct.IsValid() { buf.WriteString("DISTINCT ") - } //else if s.All.IsValid() { - // buf.WriteString("ALL ") - //} + } if s.Top.IsValid() { fmt.Fprintf(&buf, "TOP(%s) ", s.TopExpr.String()) } @@ -3810,7 +4007,6 @@ func (s *SelectStatement) String() string { buf.WriteString(window.String()) } } - // } // Write compound operator. if s.Compound != nil { @@ -3840,6 +4036,10 @@ func (s *SelectStatement) String() string { } } + if s.Limit.IsValid() { + fmt.Fprintf(&buf, "LIMIT %s", s.LimitExpr.String()) + } + return buf.String() } diff --git a/sql3/parser/ast_test.go b/sql3/parser/ast_test.go index 0c921b173..44ce35e6b 100644 --- a/sql3/parser/ast_test.go +++ b/sql3/parser/ast_test.go @@ -404,11 +404,8 @@ func TestCreateFunctionStatement_String(t *testing.T) { Type: &parser.Type{Name: &parser.Ident{Name: "int"}}, }, }, - ReturnDef: &parser.ParameterDefinition{ - Name: &parser.Variable{Name: "@scalar"}, - Type: &parser.Type{Name: &parser.Ident{Name: "int"}}, - }, - }, `CREATE FUNCTION func (@param1 int) RETURNS @scalar int AS BEGIN END`) + ReturnType: &parser.Type{Name: &parser.Ident{Name: "int"}}, + }, `CREATE FUNCTION func (@param1 int) RETURNS int AS BEGIN END`) AssertStatementStringer(t, &parser.CreateFunctionStatement{ IfNotExists: pos(0), @@ -419,11 +416,8 @@ func TestCreateFunctionStatement_String(t *testing.T) { Type: &parser.Type{Name: &parser.Ident{Name: "int"}}, }, }, - ReturnDef: &parser.ParameterDefinition{ - Name: &parser.Variable{Name: "@scalar"}, - Type: &parser.Type{Name: &parser.Ident{Name: "int"}}, - }, - }, `CREATE FUNCTION IF NOT EXISTS func (@param1 int) RETURNS @scalar int AS BEGIN END`) + ReturnType: &parser.Type{Name: &parser.Ident{Name: "int"}}, + }, `CREATE FUNCTION IF NOT EXISTS func (@param1 int) RETURNS int AS BEGIN END`) } func TestCreateViewStatement_String(t *testing.T) { @@ -866,34 +860,34 @@ func TestSelectStatement_String(t *testing.T) { }, }, `SELECT * FROM (SELECT *)`) - AssertStatementStringer(t, &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - Source: &parser.QualifiedTableName{Name: &parser.Ident{Name: "tbl"}}, - Windows: []*parser.Window{ - { - Name: &parser.Ident{Name: "win1"}, - Definition: &parser.WindowDefinition{ - Base: &parser.Ident{Name: "base"}, - Partitions: []parser.Expr{&parser.Ident{Name: "x"}, &parser.Ident{Name: "y"}}, - OrderingTerms: []*parser.OrderingTerm{ - {X: &parser.Ident{Name: "x"}, Asc: pos(0), NullsFirst: pos(0)}, - {X: &parser.Ident{Name: "y"}, Desc: pos(0), NullsLast: pos(0)}, - }, - Frame: &parser.FrameSpec{ - Range: pos(0), - UnboundedX: pos(0), - PrecedingX: pos(0), - }, - }, - }, - { - Name: &parser.Ident{Name: "win2"}, - Definition: &parser.WindowDefinition{ - Base: &parser.Ident{Name: "base2"}, - }, - }, - }, - }, `SELECT * FROM tbl WINDOW win1 AS (base PARTITION BY x, y ORDER BY x ASC NULLS FIRST, y DESC NULLS LAST RANGE UNBOUNDED PRECEDING), win2 AS (base2)`) + // AssertStatementStringer(t, &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // Source: &parser.QualifiedTableName{Name: &parser.Ident{Name: "tbl"}}, + // Windows: []*parser.Window{ + // { + // Name: &parser.Ident{Name: "win1"}, + // Definition: &parser.WindowDefinition{ + // Base: &parser.Ident{Name: "base"}, + // Partitions: []parser.Expr{&parser.Ident{Name: "x"}, &parser.Ident{Name: "y"}}, + // OrderingTerms: []*parser.OrderingTerm{ + // {X: &parser.Ident{Name: "x"}, Asc: pos(0), NullsFirst: pos(0)}, + // {X: &parser.Ident{Name: "y"}, Desc: pos(0), NullsLast: pos(0)}, + // }, + // Frame: &parser.FrameSpec{ + // Range: pos(0), + // UnboundedX: pos(0), + // PrecedingX: pos(0), + // }, + // }, + // }, + // { + // Name: &parser.Ident{Name: "win2"}, + // Definition: &parser.WindowDefinition{ + // Base: &parser.Ident{Name: "base2"}, + // }, + // }, + // }, + // }, `SELECT * FROM tbl WINDOW win1 AS (base PARTITION BY x, y ORDER BY x ASC NULLS FIRST, y DESC NULLS LAST RANGE UNBOUNDED PRECEDING), win2 AS (base2)`) // AssertStatementStringer(t, &sql.SelectStatement{ // WithClause: &sql.WithClause{ @@ -914,38 +908,38 @@ func TestSelectStatement_String(t *testing.T) { // }, // }, `WITH "cte" ("x", "y") AS (SELECT *) VALUES (1, 2), (3, 4)`) - AssertStatementStringer(t, &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - Union: pos(0), - Compound: &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - }, - }, `SELECT * UNION SELECT *`) + // AssertStatementStringer(t, &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // Union: pos(0), + // Compound: &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // }, + // }, `SELECT * UNION SELECT *`) - AssertStatementStringer(t, &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - Union: pos(0), - UnionAll: pos(0), - Compound: &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - }, - }, `SELECT * UNION ALL SELECT *`) + // AssertStatementStringer(t, &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // Union: pos(0), + // UnionAll: pos(0), + // Compound: &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // }, + // }, `SELECT * UNION ALL SELECT *`) - AssertStatementStringer(t, &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - Intersect: pos(0), - Compound: &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - }, - }, `SELECT * INTERSECT SELECT *`) + // AssertStatementStringer(t, &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // Intersect: pos(0), + // Compound: &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // }, + // }, `SELECT * INTERSECT SELECT *`) - AssertStatementStringer(t, &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - Except: pos(0), - Compound: &parser.SelectStatement{ - Columns: []*parser.ResultColumn{{Star: pos(0)}}, - }, - }, `SELECT * EXCEPT SELECT *`) + // AssertStatementStringer(t, &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // Except: pos(0), + // Compound: &parser.SelectStatement{ + // Columns: []*parser.ResultColumn{{Star: pos(0)}}, + // }, + // }, `SELECT * EXCEPT SELECT *`) AssertStatementStringer(t, &parser.SelectStatement{ Columns: []*parser.ResultColumn{{Star: pos(0)}}, diff --git a/sql3/parser/parser.go b/sql3/parser/parser.go index fe8ebea15..48e4fbcc5 100644 --- a/sql3/parser/parser.go +++ b/sql3/parser/parser.go @@ -45,10 +45,10 @@ func (p *Parser) ParseStatement() (stmt Statement, err error) { switch tok := p.peek(); tok { case EOF: return nil, io.EOF - //case EXPLAIN: - // if stmt, err = p.parseExplainStatement(); err != nil { - // return stmt, err - // } + case EXPLAIN: + if stmt, err = p.parseExplainStatement(); err != nil { + return stmt, err + } default: if stmt, err = p.parseNonExplainStatement(); err != nil { return stmt, err @@ -64,7 +64,6 @@ func (p *Parser) ParseStatement() (stmt Statement, err error) { return stmt, nil } -/* // parseExplain parses EXPLAIN [QUERY PLAN] STMT. func (p *Parser) parseExplainStatement() (_ *ExplainStatement, err error) { var tok Token @@ -89,7 +88,7 @@ func (p *Parser) parseExplainStatement() (_ *ExplainStatement, err error) { return &stmt, err } return &stmt, nil -}*/ +} // parseStmt parses all statement types. func (p *Parser) parseNonExplainStatement() (Statement, error) { @@ -102,10 +101,14 @@ func (p *Parser) parseNonExplainStatement() (Statement, error) { return p.parseBulkInsertStatement() case CREATE: return p.parseCreateStatement() + case COPY: + return p.parseCopyStatement() case DROP: return p.parseDropStatement() case SELECT: return p.parseSelectStatement(false, nil) + case PREDICT: + return p.parsePredictStatement() case INSERT, REPLACE: return p.parseInsertStatement(nil) case UPDATE: @@ -337,8 +340,10 @@ func (p *Parser) parseCreateStatement() (Statement, error) { return p.parseCreateIndexStatement(pos)*/ case FUNCTION: return p.parseCreateFunctionStatement(pos) + case MODEL: + return p.parseCreateModelStatement(pos) default: - return nil, p.errorExpected(pos, tok, "DATABASE, TABLE, VIEW or FUNCTION") + return nil, p.errorExpected(pos, tok, "DATABASE, TABLE, VIEW, FUNCTION or MODEL") } } @@ -373,6 +378,8 @@ func (p *Parser) parseDropStatement() (Statement, error) { return p.parseDropIndexStatement(pos)*/ case FUNCTION: return p.parseDropFunctionStatement(pos) + case MODEL: + return p.parseDropModelStatement(pos) default: return nil, p.errorExpected(pos, tok, "DATABASE, TABLE, VIEW or FUNCTION") } @@ -1158,6 +1165,72 @@ func (p *Parser) parseDropTableStatement(dropPos Pos) (_ *DropTableStatement, er return &stmt, nil } +func (p *Parser) parseCopyStatement() (_ *CopyStatement, err error) { + assert(p.peek() == COPY) + + var stmt CopyStatement + stmt.Copy, _, _ = p.scan() + + ident, err := p.parseIdent("table name") + if err != nil { + return &stmt, err + } + + if stmt.Source, err = p.parseQualifiedTableName(ident); err != nil { + return &stmt, err + } + + if p.peek() != TO { + return &stmt, p.errorExpected(p.pos, p.tok, "TO") + } + stmt.To, _, _ = p.scan() + + if stmt.TargetName, err = p.parseIdent("table name"); err != nil { + return &stmt, err + } + + // parse optional "WHERE expr" + if p.peek() == WHERE { + stmt.Where, _, _ = p.scan() + if stmt.WhereExpr, err = p.ParseExpr(); err != nil { + return &stmt, err + } + } + + // options + if p.peek() == WITH { + stmt.With, _, _ = p.scan() + if !isCopyOptionStartToken(p.peek(), p) { + return &stmt, p.errorExpected(p.pos, p.tok, "URL or APIKEY") + } + + for { + option, err := p.parseIdent("copy option") + if err != nil { + return &stmt, err + } + + switch strings.ToLower(option.Name) { + case "url": + stmt.Url, err = p.ParseExpr() + if err != nil { + return &stmt, err + } + + case "apikey": + stmt.ApiKey, err = p.ParseExpr() + if err != nil { + return &stmt, err + } + } + if !isCopyOptionStartToken(p.peek(), p) { + break + } + } + } + return &stmt, nil +} + func (p *Parser) parseCreateViewStatement(createPos Pos) (_ *CreateViewStatement, err error) { assert(p.peek() == VIEW) @@ -1286,6 +1359,29 @@ func (p *Parser) parseDropViewStatement(dropPos Pos) (_ *DropViewStatement, err return &stmt, nil } +func (p *Parser) parseDropModelStatement(dropPos Pos) (_ *DropModelStatement, err error) { + assert(p.peek() == MODEL) + + var stmt DropModelStatement + stmt.Drop = dropPos + stmt.Model, _, _ = p.scan() + + // Parse optional "IF EXISTS". + if p.peek() == IF { + stmt.If, _, _ = p.scan() + if p.peek() != EXISTS { + return &stmt, p.errorExpected(p.pos, p.tok, "EXISTS") + } + stmt.IfExists, _, _ = p.scan() + } + + if stmt.Name, err = p.parseIdent("view name"); err != nil { + return &stmt, err + } + + return &stmt, nil +} + /*func (p *Parser) parseCreateIndexStatement(createPos Pos) (_ *CreateIndexStatement, err error) { assert(p.peek() == INDEX || p.peek() == UNIQUE) @@ -1457,11 +1553,35 @@ func (p *Parser) parseCreateFunctionStatement(createPos Pos) (_ *CreateFunctionS return &stmt, p.errorExpected(p.pos, p.tok, "RETURNS") } stmt.Returns, _, _ = p.scan() - stmt.ReturnDef, err = p.parseParameterDefinition() + stmt.ReturnType, err = p.parseType() if err != nil { return &stmt, err } + // options + if p.peek() == WITH { + stmt.With, _, _ = p.scan() + stmt.Options = make([]*FunctionOptionDefinition, 0) + for { + option, err := p.parseIdent("function option") + if err != nil { + return &stmt, err + } + + expr, err := p.ParseExpr() + if err != nil { + return &stmt, err + } + stmt.Options = append(stmt.Options, &FunctionOptionDefinition{ + Name: option, + OptionExpr: expr, + }) + if p.peek() == AS { + break + } + } + } + if p.peek() != AS { return &stmt, p.errorExpected(p.pos, p.tok, "AS") } @@ -1473,7 +1593,7 @@ func (p *Parser) parseCreateFunctionStatement(createPos Pos) (_ *CreateFunctionS stmt.Begin, _, _ = p.scan() for { - s, err := p.parseFunctionBodyStatement() + s, err := p.parseFunctionBodyStatement(&stmt) if err != nil { return &stmt, err } @@ -1494,8 +1614,15 @@ func (p *Parser) parseCreateFunctionStatement(createPos Pos) (_ *CreateFunctionS return &stmt, nil } -func (p *Parser) parseFunctionBodyStatement() (stmt Statement, err error) { +func (p *Parser) parseFunctionBodyStatement(cf *CreateFunctionStatement) (stmt Statement, err error) { switch p.peek() { + case RETURN: + s, err := p.parseReturnStatement() + if err != nil { + return stmt, err + } + cf.Body = append(cf.Body, s) + case END: break default: @@ -1507,6 +1634,20 @@ func (p *Parser) parseFunctionBodyStatement() (stmt Statement, err error) { return stmt, nil } +func (p *Parser) parseReturnStatement() (_ *ReturnStatement, err error) { + assert(p.peek() == RETURN) + + var stmt ReturnStatement + stmt.Return, _, _ = p.scan() + + expr, err := p.ParseExpr() + if err != nil { + return &stmt, err + } + stmt.ReturnExpr = expr + return &stmt, nil +} + func (p *Parser) parseDropFunctionStatement(dropPos Pos) (_ *DropFunctionStatement, err error) { assert(p.peek() == FUNCTION) @@ -1530,6 +1671,69 @@ func (p *Parser) parseDropFunctionStatement(dropPos Pos) (_ *DropFunctionStateme return &stmt, nil } +func (p *Parser) parseCreateModelStatement(createPos Pos) (_ *CreateModelStatement, err error) { + assert(p.peek() == MODEL) + + var stmt CreateModelStatement + stmt.Create = createPos + stmt.Model, _, _ = p.scan() + + // Parse optional "IF NOT EXISTS". + if p.peek() == IF { + stmt.If, _, _ = p.scan() + + if p.peek() != NOT { + return &stmt, p.errorExpected(p.pos, p.tok, "NOT") + } + stmt.IfNot, _, _ = p.scan() + + if p.peek() != EXISTS { + return &stmt, p.errorExpected(p.pos, p.tok, "EXISTS") + } + stmt.IfNotExists, _, _ = p.scan() + } + + if stmt.Name, err = p.parseIdent("model name"); err != nil { + return &stmt, err + } + + // options + if p.peek() != WITH { + return &stmt, p.errorExpected(p.pos, p.tok, "WITH") + } + stmt.With, _, _ = p.scan() + + stmt.Options = make([]*ModelOptionDefinition, 0) + for { + option, err := p.parseIdent("model option") + if err != nil { + return &stmt, err + } + + expr, err := p.ParseExpr() + if err != nil { + return &stmt, err + } + stmt.Options = append(stmt.Options, &ModelOptionDefinition{ + Name: option, + OptionExpr: expr, + }) + if p.peek() == AS { + break + } + } + + if p.peek() != AS { + return &stmt, p.errorExpected(p.pos, p.tok, "AS") + } + stmt.As, _, _ = p.scan() + + if stmt.ModelQuery, err = p.parseSelectStatement(false, nil); err != nil { + return &stmt, err + } + return &stmt, nil +} + func (p *Parser) parseIdent(desc string) (*Ident, error) { pos, tok, lit := p.scan() switch tok { @@ -2152,82 +2356,87 @@ func (p *Parser) parseSelectStatement(compounded bool, withClause *WithClause) ( // } //} - switch p.peek() { - /*case VALUES: - stmt.Values, _, _ = p.scan() + if p.peek() != SELECT { + return &stmt, p.errorExpected(p.pos, p.tok, "SELECT") + } - for { - var list ExprList - if p.peek() != LP { - return &stmt, p.errorExpected(p.pos, p.tok, "left paren") + stmt.Select, _, _ = p.scan() + + // Parse optional "DISTINCT". + if tok := p.peek(); tok == DISTINCT { + stmt.Distinct, _, _ = p.scan() + } + + if p.peek() == TOP { + stmt.Top, _, _ = p.scan() + if p.peek() == LP { + _, _, _ = p.scan() } - list.Lparen, _, _ = p.scan() + if stmt.TopExpr, err = p.ParseExpr(); err != nil { + return &stmt, err + } + if p.peek() == RP { + _, _, _ = p.scan() + } + } + + if p.peek() == TOPN { + stmt.TopN, _, _ = p.scan() + if p.peek() == LP { + _, _, _ = p.scan() + } + if stmt.TopExpr, err = p.ParseExpr(); err != nil { + return &stmt, err + } + if p.peek() == RP { + _, _, _ = p.scan() + } + } + + // Parse result columns. + for { + col, err := p.parseResultColumn() + if err != nil { + return &stmt, err + } + stmt.Columns = append(stmt.Columns, col) + + if p.peek() != COMMA { + break + } + p.scan() + } + + // Parse FROM clause. + if p.peek() == FROM { + stmt.From, _, _ = p.scan() + if stmt.Source, err = p.parseSource(); err != nil { + return &stmt, err + } + } + + // Parse WHERE clause. + if p.peek() == WHERE { + stmt.Where, _, _ = p.scan() + if stmt.WhereExpr, err = p.ParseExpr(); err != nil { + return &stmt, err + } + } + + // Parse GROUP BY/HAVING clause. + if p.peek() == GROUP { + stmt.Group, _, _ = p.scan() + if p.peek() != BY { + return &stmt, p.errorExpected(p.pos, p.tok, "BY") + } + stmt.GroupBy, _, _ = p.scan() for { expr, err := p.ParseExpr() if err != nil { return &stmt, err } - list.Exprs = append(list.Exprs, expr) - - if p.peek() == RP { - break - } else if p.peek() != COMMA { - return &stmt, p.errorExpected(p.pos, p.tok, "comma or right paren") - } - p.scan() - } - list.Rparen, _, _ = p.scan() - stmt.ValueLists = append(stmt.ValueLists, &list) - - if p.peek() != COMMA { - break - } - p.scan() - - }*/ - - case SELECT: - stmt.Select, _, _ = p.scan() - - // Parse optional "DISTINCT". - if tok := p.peek(); tok == DISTINCT { - stmt.Distinct, _, _ = p.scan() - } - - if p.peek() == TOP { - stmt.Top, _, _ = p.scan() - if p.peek() == LP { - _, _, _ = p.scan() - } - if stmt.TopExpr, err = p.ParseExpr(); err != nil { - return &stmt, err - } - if p.peek() == RP { - _, _, _ = p.scan() - } - } - - if p.peek() == TOPN { - stmt.TopN, _, _ = p.scan() - if p.peek() == LP { - _, _, _ = p.scan() - } - if stmt.TopExpr, err = p.ParseExpr(); err != nil { - return &stmt, err - } - if p.peek() == RP { - _, _, _ = p.scan() - } - } - - // Parse result columns. - for { - col, err := p.parseResultColumn() - if err != nil { - return &stmt, err - } - stmt.Columns = append(stmt.Columns, col) + stmt.GroupByExprs = append(stmt.GroupByExprs, expr) if p.peek() != COMMA { break @@ -2235,101 +2444,61 @@ func (p *Parser) parseSelectStatement(compounded bool, withClause *WithClause) ( p.scan() } - // Parse FROM clause. - if p.peek() == FROM { - stmt.From, _, _ = p.scan() - if stmt.Source, err = p.parseSource(); err != nil { + // Parse optional HAVING clause. + if p.peek() == HAVING { + stmt.Having, _, _ = p.scan() + if stmt.HavingExpr, err = p.ParseExpr(); err != nil { return &stmt, err } } - - // Parse WHERE clause. - if p.peek() == WHERE { - stmt.Where, _, _ = p.scan() - if stmt.WhereExpr, err = p.ParseExpr(); err != nil { - return &stmt, err - } - } - - // Parse GROUP BY/HAVING clause. - if p.peek() == GROUP { - stmt.Group, _, _ = p.scan() - if p.peek() != BY { - return &stmt, p.errorExpected(p.pos, p.tok, "BY") - } - stmt.GroupBy, _, _ = p.scan() - - for { - expr, err := p.ParseExpr() - if err != nil { - return &stmt, err - } - stmt.GroupByExprs = append(stmt.GroupByExprs, expr) - - if p.peek() != COMMA { - break - } - p.scan() - } - - // Parse optional HAVING clause. - if p.peek() == HAVING { - stmt.Having, _, _ = p.scan() - if stmt.HavingExpr, err = p.ParseExpr(); err != nil { - return &stmt, err - } - } - } - - // Parse WINDOW clause. - if p.peek() == WINDOW { - stmt.Window, _, _ = p.scan() - - for { - var window Window - if window.Name, err = p.parseIdent("window name"); err != nil { - return &stmt, err - } - - if p.peek() != AS { - return &stmt, p.errorExpected(p.pos, p.tok, "AS") - } - window.As, _, _ = p.scan() - - if window.Definition, err = p.parseWindowDefinition(); err != nil { - return &stmt, err - } - - stmt.Windows = append(stmt.Windows, &window) - - if p.peek() != COMMA { - break - } - p.scan() - } - } - default: - return &stmt, p.errorExpected(p.pos, p.tok, "SELECT") } + // Parse WINDOW clause. + // if p.peek() == WINDOW { + // stmt.Window, _, _ = p.scan() + + // for { + // var window Window + // if window.Name, err = p.parseIdent("window name"); err != nil { + // return &stmt, err + // } + + // if p.peek() != AS { + // return &stmt, p.errorExpected(p.pos, p.tok, "AS") + // } + // window.As, _, _ = p.scan() + + // if window.Definition, err = p.parseWindowDefinition(); err != nil { + // return &stmt, err + // } + + // stmt.Windows = append(stmt.Windows, &window) + + // if p.peek() != COMMA { + // break + // } + // p.scan() + // } + // } + // Optionally compound additional SELECT/VALUES. - switch tok := p.peek(); tok { - case UNION, INTERSECT, EXCEPT: - if tok == UNION { - stmt.Union, _, _ = p.scan() - if p.peek() == ALL { - stmt.UnionAll, _, _ = p.scan() - } - } else if tok == INTERSECT { - stmt.Intersect, _, _ = p.scan() - } else { - stmt.Except, _, _ = p.scan() - } + // switch tok := p.peek(); tok { + // case UNION, INTERSECT, EXCEPT: + // if tok == UNION { + // stmt.Union, _, _ = p.scan() + // if p.peek() == ALL { + // stmt.UnionAll, _, _ = p.scan() + // } + // } else if tok == INTERSECT { + // stmt.Intersect, _, _ = p.scan() + // } else { + // stmt.Except, _, _ = p.scan() + // } - if stmt.Compound, err = p.parseSelectStatement(true, nil); err != nil { - return &stmt, err - } - } + // if stmt.Compound, err = p.parseSelectStatement(true, nil); err != nil { + // return &stmt, err + // } + // } // Parse ORDER BY clause. if !compounded && p.peek() == ORDER { @@ -2353,6 +2522,13 @@ func (p *Parser) parseSelectStatement(compounded bool, withClause *WithClause) ( } } + // Parse LIMIT clause. + if !compounded && p.peek() == LIMIT { + stmt.Limit, _, _ = p.scan() + if stmt.LimitExpr, err = p.ParseExpr(); err != nil { + return &stmt, err + } + } return &stmt, nil } @@ -2704,6 +2880,26 @@ func (p *Parser) parseTableValuedFunction(ident *Ident) (_ *TableValuedFunction, return &cte, nil }*/ +func (p *Parser) parsePredictStatement() (_ *PredictStatement, err error) { + assert(p.peek() == PREDICT) + + var stmt PredictStatement + stmt.Predict, _, _ = p.scan() + if p.peek() != USING { + return &stmt, p.errorExpected(p.pos, p.tok, "USING") + } + stmt.Using, _, _ = p.scan() + + if stmt.ModelName, err = p.parseIdent("model name"); err != nil { + return &stmt, err + } + + if stmt.InputQuery, err = p.parseSelectStatement(false, nil); err != nil { + return &stmt, err + } + return &stmt, nil +} + func (p *Parser) mustParseLiteral() Expr { assert(isLiteralToken(p.tok)) pos, tok, lit := p.scan() @@ -2742,9 +2938,13 @@ func (p *Parser) parseOperand() (expr Expr, err error) { case VARIABLE: return &Variable{Name: lit, NamePos: pos}, nil case MIN, MAX: - ident := &Ident{Name: lit, NamePos: pos, Quoted: tok == QIDENT} - return p.parseCall(ident) - case STRING: + pk := p.peek() + if pk == LP { + ident := &Ident{Name: lit, NamePos: pos, Quoted: false} + return p.parseCall(ident) + } + return nil, p.errorExpected(p.pos, pk, "call expression") + case STRING, BLOB: return &StringLit{ValuePos: pos, Value: lit}, nil case FLOAT: return &FloatLit{ValuePos: pos, Value: lit}, nil @@ -3629,6 +3829,22 @@ func isBulkInsertOptionStartToken(tok Token, p *Parser) bool { return false } +func isCopyOptionStartToken(tok Token, p *Parser) bool { + switch tok { + case IDENT: + ident, err := p.parseIdent("copy option") + defer p.unscan() + if err != nil { + return false + } + switch strings.ToUpper(ident.Name) { + case "URL", "APIKEY": + return true + } + } + return false +} + // isConstraintStartToken returns true if tok is the initial token of a constraint. func isConstraintStartToken(tok Token, isTable bool) bool { switch tok { diff --git a/sql3/parser/parser_test.go b/sql3/parser/parser_test.go index deb53a5f9..72b316b52 100644 --- a/sql3/parser/parser_test.go +++ b/sql3/parser/parser_test.go @@ -40,6 +40,9 @@ func TestParser_ParseMinMaxColumnConstraints(t *testing.T) { t.Run("ErrNoKey", func(t *testing.T) { AssertParseStatementError(t, `CREATE TABLE tbl (col1 INT MIN`, `1:30: expected expression, found 'EOF'`) }) + t.Run("ErrNoCall", func(t *testing.T) { + AssertParseStatementError(t, `SELECT MIN;`, `1:11: expected call expression, found ';'`) + }) t.Run("Simple", func(t *testing.T) { AssertParseStatement(t, `CREATE TABLE tbl (col1 INT MIN 0)`, &parser.CreateTableStatement{ Create: pos(0), @@ -469,7 +472,7 @@ func TestParser_ParseAlterStatement(t *testing.T) { func TestParser_ParseFunctionStatement(t *testing.T) { t.Run("CreateFunction", func(t *testing.T) { - AssertParseStatement(t, `CREATE FUNCTION IF NOT EXISTS func (@param1 int, @param2 string) returns @scalar int as begin end`, &parser.CreateFunctionStatement{ + AssertParseStatement(t, `CREATE FUNCTION IF NOT EXISTS func (@param1 int, @param2 string) returns int as begin end`, &parser.CreateFunctionStatement{ Create: pos(0), Function: pos(7), If: pos(16), @@ -487,15 +490,12 @@ func TestParser_ParseFunctionStatement(t *testing.T) { Type: &parser.Type{Name: &parser.Ident{NamePos: pos(57), Name: "string"}}, }, }, - Rparen: pos(63), - Returns: pos(65), - ReturnDef: &parser.ParameterDefinition{ - Name: &parser.Variable{Name: "@scalar", NamePos: pos(73)}, - Type: &parser.Type{Name: &parser.Ident{NamePos: pos(81), Name: "int"}}, - }, - As: pos(85), - Begin: pos(88), - End: pos(94), + Rparen: pos(63), + Returns: pos(65), + ReturnType: &parser.Type{Name: &parser.Ident{NamePos: pos(73), Name: "int"}}, + As: pos(77), + Begin: pos(80), + End: pos(86), }) // AssertParseStatement(t, `CREATE TRIGGER IF NOT EXISTS trig BEFORE INSERT ON tbl BEGIN DELETE FROM new; END`, &parser.CreateFunctionStatement{ // Create: pos(0), @@ -952,7 +952,7 @@ func TestParser_ParseStatement(t *testing.T) { }, }) - AssertParseStatementError(t, `CREATE`, `1:1: expected DATABASE, TABLE, VIEW or FUNCTION`) + AssertParseStatementError(t, `CREATE`, `1:1: expected DATABASE, TABLE, VIEW, FUNCTION or MODEL`) AssertParseStatementError(t, `CREATE DATABASE`, `1:15: expected database name, found 'EOF'`) AssertParseStatementError(t, `CREATE DATABASE IF`, `1:18: expected NOT, found 'EOF'`) AssertParseStatementError(t, `CREATE DATABASE IF NOT`, `1:22: expected EXISTS, found 'EOF'`) @@ -2212,29 +2212,29 @@ func TestParser_ParseStatement(t *testing.T) { Having: pos(22), HavingExpr: &parser.BoolLit{ValuePos: pos(29), Value: true}, }) - AssertParseStatement(t, `SELECT * WINDOW win1 AS (), win2 AS ()`, &parser.SelectStatement{ - Select: pos(0), - Columns: []*parser.ResultColumn{{Star: pos(7)}}, - Window: pos(9), - Windows: []*parser.Window{ - { - Name: &parser.Ident{NamePos: pos(16), Name: "win1"}, - As: pos(21), - Definition: &parser.WindowDefinition{ - Lparen: pos(24), - Rparen: pos(25), - }, - }, - { - Name: &parser.Ident{NamePos: pos(28), Name: "win2"}, - As: pos(33), - Definition: &parser.WindowDefinition{ - Lparen: pos(36), - Rparen: pos(37), - }, - }, - }, - }) + // AssertParseStatement(t, `SELECT * WINDOW win1 AS (), win2 AS ()`, &parser.SelectStatement{ + // Select: pos(0), + // Columns: []*parser.ResultColumn{{Star: pos(7)}}, + // Window: pos(9), + // Windows: []*parser.Window{ + // { + // Name: &parser.Ident{NamePos: pos(16), Name: "win1"}, + // As: pos(21), + // Definition: &parser.WindowDefinition{ + // Lparen: pos(24), + // Rparen: pos(25), + // }, + // }, + // { + // Name: &parser.Ident{NamePos: pos(28), Name: "win2"}, + // As: pos(33), + // Definition: &parser.WindowDefinition{ + // Lparen: pos(36), + // Rparen: pos(37), + // }, + // }, + // }, + // }) AssertParseStatement(t, `SELECT * ORDER BY foo ASC, bar DESC`, &parser.SelectStatement{ Select: pos(0), @@ -2249,64 +2249,64 @@ func TestParser_ParseStatement(t *testing.T) { }, }) - AssertParseStatement(t, `SELECT * UNION SELECT * ORDER BY foo`, &parser.SelectStatement{ - Select: pos(0), - Columns: []*parser.ResultColumn{ - {Star: pos(7)}, - }, - Union: pos(9), - Compound: &parser.SelectStatement{ - Select: pos(15), - Columns: []*parser.ResultColumn{ - {Star: pos(22)}, - }, - }, - Order: pos(24), - OrderBy: pos(30), - OrderingTerms: []*parser.OrderingTerm{ - {X: &parser.Ident{NamePos: pos(33), Name: "foo"}}, - }, - }) - AssertParseStatement(t, `SELECT * UNION ALL SELECT *`, &parser.SelectStatement{ - Select: pos(0), - Columns: []*parser.ResultColumn{ - {Star: pos(7)}, - }, - Union: pos(9), - UnionAll: pos(15), - Compound: &parser.SelectStatement{ - Select: pos(19), - Columns: []*parser.ResultColumn{ - {Star: pos(26)}, - }, - }, - }) - AssertParseStatement(t, `SELECT * INTERSECT SELECT *`, &parser.SelectStatement{ - Select: pos(0), - Columns: []*parser.ResultColumn{ - {Star: pos(7)}, - }, - Intersect: pos(9), - Compound: &parser.SelectStatement{ - Select: pos(19), - Columns: []*parser.ResultColumn{ - {Star: pos(26)}, - }, - }, - }) - AssertParseStatement(t, `SELECT * EXCEPT SELECT *`, &parser.SelectStatement{ - Select: pos(0), - Columns: []*parser.ResultColumn{ - {Star: pos(7)}, - }, - Except: pos(9), - Compound: &parser.SelectStatement{ - Select: pos(16), - Columns: []*parser.ResultColumn{ - {Star: pos(23)}, - }, - }, - }) + // AssertParseStatement(t, `SELECT * UNION SELECT * ORDER BY foo`, &parser.SelectStatement{ + // Select: pos(0), + // Columns: []*parser.ResultColumn{ + // {Star: pos(7)}, + // }, + // Union: pos(9), + // Compound: &parser.SelectStatement{ + // Select: pos(15), + // Columns: []*parser.ResultColumn{ + // {Star: pos(22)}, + // }, + // }, + // Order: pos(24), + // OrderBy: pos(30), + // OrderingTerms: []*parser.OrderingTerm{ + // {X: &parser.Ident{NamePos: pos(33), Name: "foo"}}, + // }, + // }) + // AssertParseStatement(t, `SELECT * UNION ALL SELECT *`, &parser.SelectStatement{ + // Select: pos(0), + // Columns: []*parser.ResultColumn{ + // {Star: pos(7)}, + // }, + // Union: pos(9), + // UnionAll: pos(15), + // Compound: &parser.SelectStatement{ + // Select: pos(19), + // Columns: []*parser.ResultColumn{ + // {Star: pos(26)}, + // }, + // }, + // }) + // AssertParseStatement(t, `SELECT * INTERSECT SELECT *`, &parser.SelectStatement{ + // Select: pos(0), + // Columns: []*parser.ResultColumn{ + // {Star: pos(7)}, + // }, + // Intersect: pos(9), + // Compound: &parser.SelectStatement{ + // Select: pos(19), + // Columns: []*parser.ResultColumn{ + // {Star: pos(26)}, + // }, + // }, + // }) + // AssertParseStatement(t, `SELECT * EXCEPT SELECT *`, &parser.SelectStatement{ + // Select: pos(0), + // Columns: []*parser.ResultColumn{ + // {Star: pos(7)}, + // }, + // Except: pos(9), + // Compound: &parser.SelectStatement{ + // Select: pos(16), + // Columns: []*parser.ResultColumn{ + // {Star: pos(23)}, + // }, + // }, + // }) /*AssertParseStatement(t, `VALUES (1, 2), (3, 4)`, &parser.SelectStatement{ Values: pos(0), @@ -3399,11 +3399,11 @@ func TestParser_ParseStatement(t *testing.T) { AssertParseStatementError(t, `SELECT * GROUP BY`, `1:17: expected expression, found 'EOF'`) AssertParseStatementError(t, `SELECT * GROUP BY foo bar`, `1:23: expected semicolon or EOF, found bar`) AssertParseStatementError(t, `SELECT * GROUP BY foo HAVING`, `1:28: expected expression, found 'EOF'`) - AssertParseStatementError(t, `SELECT * WINDOW`, `1:15: expected window name, found 'EOF'`) - AssertParseStatementError(t, `SELECT * WINDOW win1`, `1:20: expected AS, found 'EOF'`) - AssertParseStatementError(t, `SELECT * WINDOW win1 AS`, `1:23: expected left paren, found 'EOF'`) - AssertParseStatementError(t, `SELECT * WINDOW win1 AS (`, `1:25: expected right paren, found 'EOF'`) - AssertParseStatementError(t, `SELECT * WINDOW win1 AS () win2`, `1:28: expected semicolon or EOF, found win2`) + // AssertParseStatementError(t, `SELECT * WINDOW`, `1:15: expected window name, found 'EOF'`) + // AssertParseStatementError(t, `SELECT * WINDOW win1`, `1:20: expected AS, found 'EOF'`) + // AssertParseStatementError(t, `SELECT * WINDOW win1 AS`, `1:23: expected left paren, found 'EOF'`) + // AssertParseStatementError(t, `SELECT * WINDOW win1 AS (`, `1:25: expected right paren, found 'EOF'`) + // AssertParseStatementError(t, `SELECT * WINDOW win1 AS () win2`, `1:28: expected semicolon or EOF, found win2`) AssertParseStatementError(t, `SELECT * ORDER`, `1:14: expected BY, found 'EOF'`) AssertParseStatementError(t, `SELECT * ORDER BY`, `1:17: expected expression, found 'EOF'`) AssertParseStatementError(t, `SELECT * ORDER BY 1,`, `1:20: expected expression, found 'EOF'`) diff --git a/sql3/parser/token.go b/sql3/parser/token.go index ef16f7ef2..165e07243 100644 --- a/sql3/parser/token.go +++ b/sql3/parser/token.go @@ -78,8 +78,6 @@ const ( ACTION ADD AFTER - AGG_COLUMN - AGG_FUNCTION ALL ALTER ANALYZE @@ -106,9 +104,9 @@ const ( COMMENT CONFLICT CONSTRAINT + COPY CREATE CROSS - CTIME_KW CURRENT CURRENT_DATE CURRENT_TIMESTAMP @@ -167,11 +165,13 @@ const ( LAST LEFT LIKE + LIMIT LRU MAP MATCH MAX MIN + MODEL NO NOT NOTBETWEEN @@ -194,6 +194,7 @@ const ( PLAN PRAGMA PRECEDING + PREDICT PRIMARY QUERY RANGE @@ -304,8 +305,6 @@ var tokens = [...]string{ ACTION: "ACTION", ADD: "ADD", AFTER: "AFTER", - AGG_COLUMN: "AGG_COLUMN", - AGG_FUNCTION: "AGG_FUNCTION", ALL: "ALL", ALTER: "ALTER", ANALYZE: "ANALYZE", @@ -332,9 +331,9 @@ var tokens = [...]string{ COMMENT: "COMMENT", CONFLICT: "CONFLICT", CONSTRAINT: "CONSTRAINT", + COPY: "COPY", CREATE: "CREATE", CROSS: "CROSS", - CTIME_KW: "CTIME_KW", CURRENT: "CURRENT", CURRENT_DATE: "CURRENT_DATE", CURRENT_TIMESTAMP: "CURRENT_TIMESTAMP", @@ -393,11 +392,13 @@ var tokens = [...]string{ LAST: "LAST", LEFT: "LEFT", LIKE: "LIKE", + LIMIT: "LIMIT", MAP: "MAP", LRU: "LRU", MATCH: "MATCH", MAX: "MAX", MIN: "MIN", + MODEL: "MODEL", NO: "NO", NOT: "NOT", NOTBETWEEN: "NOTBETWEEN", @@ -420,6 +421,7 @@ var tokens = [...]string{ PLAN: "PLAN", PRAGMA: "PRAGMA", PRECEDING: "PRECEDING", + PREDICT: "PREDICT", PRIMARY: "PRIMARY", QUERY: "QUERY", RANGE: "RANGE", diff --git a/sql3/planner/compilecopy.go b/sql3/planner/compilecopy.go new file mode 100644 index 000000000..79d8c78bc --- /dev/null +++ b/sql3/planner/compilecopy.go @@ -0,0 +1,120 @@ +// Copyright 2022 Molecula Corp. All rights reserved. + +package planner + +import ( + "context" + + "github.com/featurebasedb/featurebase/v3/dax" + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// compileCopyStatement compiles a parser.CopyStatement AST into a PlanOperator +func (p *ExecutionPlanner) compileCopyStatement(stmt *parser.CopyStatement) (types.PlanOperator, error) { + query := NewPlanOpQuery(p, NewPlanOpNullTable(), p.sql) + query.AddWarning("🦖 here there be dragons! COPY statement is experimental.") + + // handle projections + projections := make([]types.PlanExpression, 0) + for _, c := range stmt.Source.PossibleOutputColumns() { + expr := &parser.QualifiedRef{ + Table: &parser.Ident{Name: c.TableName}, + Column: &parser.Ident{Name: c.ColumnName}, + ColumnIndex: c.ColumnIndex, + RefDataType: c.Datatype, + } + planExpr, err := p.compileExpr(expr) + if err != nil { + return nil, err + } + projections = append(projections, planExpr) + } + + // handle the where clause + where, err := p.compileExpr(stmt.WhereExpr) + if err != nil { + return nil, err + } + + // compile source + source, err := p.compileSource(query, stmt.Source) + if err != nil { + return nil, err + } + + // if we did have a where, insert the filter op + if where != nil { + source = NewPlanOpFilter(p, where, source) + } + + var compiledOp types.PlanOperator + url := "" + apiKey := "" + + if stmt.Url != nil { + lit, ok := stmt.Url.(*parser.StringLit) + if !ok { + return nil, sql3.NewErrStringLiteral(stmt.Url.Pos().Line, stmt.Url.Pos().Column) + } + url = lit.Value + } + + if stmt.ApiKey != nil { + lit, ok := stmt.ApiKey.(*parser.StringLit) + if !ok { + return nil, sql3.NewErrStringLiteral(stmt.ApiKey.Pos().Line, stmt.ApiKey.Pos().Column) + } + apiKey = lit.Value + } + + // get the source table + tname := dax.TableName(stmt.Source.String()) + tbl, err := p.schemaAPI.TableByName(context.Background(), tname) + if err != nil { + if isTableNotFoundError(err) { + return nil, sql3.NewErrTableNotFound(0, 0, stmt.Source.String()) + } + return nil, err + } + // get the ddl of source table and subst target table name + ddl := generateTableDDL(tbl, stmt.TargetName.Name) + compiledOp = NewPlanOpCopy(p, stmt.TargetName.Name, url, apiKey, ddl, NewPlanOpProjection(projections, source)) + children := []types.PlanOperator{ + compiledOp, + } + return query.WithChildren(children...) +} + +func (p *ExecutionPlanner) analyzeCopyStatement(ctx context.Context, stmt *parser.CopyStatement) error { + + // analyze source + var err error + source, err := p.analyzeSource(ctx, stmt.Source, stmt) + if err != nil { + return err + } + stmt.Source = source + + // analyze where + expr, err := p.analyzeExpression(ctx, stmt.WhereExpr, stmt) + if err != nil { + return err + } + stmt.WhereExpr = expr + + expr, err = p.analyzeExpression(ctx, stmt.Url, stmt) + if err != nil { + return err + } + stmt.Url = expr + + expr, err = p.analyzeExpression(ctx, stmt.ApiKey, stmt) + if err != nil { + return err + } + stmt.ApiKey = expr + + return nil +} diff --git a/sql3/planner/compilecreatefunction.go b/sql3/planner/compilecreatefunction.go new file mode 100644 index 000000000..fbb055b7f --- /dev/null +++ b/sql3/planner/compilecreatefunction.go @@ -0,0 +1,79 @@ +// Copyright 2023 Molecula Corp. All rights reserved. + +package planner + +import ( + "strings" + + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// compileCreateFunctionStatement compiles a parser.CreateFunctionStatement AST into a PlanOperator +func (p *ExecutionPlanner) compileCreateFunctionStatement(stmt *parser.CreateFunctionStatement) (types.PlanOperator, error) { + functionName := parser.IdentName(stmt.Name) + function := &functionSystemObject{ + name: functionName, + } + + lang := "sql" + if len(stmt.Options) > 0 { + for _, o := range stmt.Options { + switch strings.ToLower(o.Name.String()) { + case "language": + lit, ok := o.OptionExpr.(*parser.StringLit) + if !ok { + return nil, sql3.NewErrStringLiteral(o.OptionExpr.Pos().Line, o.OptionExpr.Pos().Column) + } + l := strings.ToLower(lit.Value) + switch l { + case "python": + lang = l + default: + return nil, sql3.NewErrInternalf("unsupported language '%s'", l) + } + } + } + } + function.language = lang + + // TODO(pok) - hobble user defined functions for now + + switch lang { + case "sql": + return nil, sql3.NewErrInternalf("unsupported language '%s'", lang) + case "python": + // return nil, sql3.NewErrInternalf("unsupported language '%s'", lang) + + // function body is in the return statement + if len(stmt.Body) != 1 { + return nil, sql3.NewErrInternalf("unexpected body len '%d'", len(stmt.Body)) + } + + rs, ok := stmt.Body[0].(*parser.ReturnStatement) + if !ok { + return nil, sql3.NewErrInternalf("unexpected statement type '%T'", stmt.Body[0]) + } + + bexpr, ok := rs.ReturnExpr.(*parser.StringLit) + if !ok { + return nil, sql3.NewErrInternalf("unexpected expression type '%T'", rs.ReturnExpr) + } + + function.body = bexpr.Value + default: + return nil, sql3.NewErrInternalf("unsupported language '%s'", lang) + } + + fn := NewPlanOpCreateFunction(p, stmt.IfNotExists.IsValid(), function) + fn.AddWarning("🦖 here there be dragons! CREATE FUNCTION statement is experimental.") + + query := NewPlanOpQuery(p, fn, p.sql) + return query, nil +} + +func (p *ExecutionPlanner) analyzeCreateFunctionStatement(stmt *parser.CreateFunctionStatement) error { + + return nil +} diff --git a/sql3/planner/compilecreatemodel.go b/sql3/planner/compilecreatemodel.go new file mode 100644 index 000000000..9cd57a056 --- /dev/null +++ b/sql3/planner/compilecreatemodel.go @@ -0,0 +1,180 @@ +// Copyright 2022 Molecula Corp. All rights reserved. + +package planner + +import ( + "context" + "strings" + + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// TODO (pok) what does 'if not exists' do? + +// compileCreateModelStatement compiles a parser.CreateModelStatement AST into a PlanOperator +func (p *ExecutionPlanner) compileCreateModelStatement(stmt *parser.CreateModelStatement) (types.PlanOperator, error) { + modelName := parser.IdentName(stmt.Name) + + // does the model exist + obj, err := p.getModelByName(modelName) + if err != nil { + return nil, err + } + if obj != nil { + return nil, sql3.NewErrInternalf("model '%s' already exists", modelName) + } + + // if we got to here model does not exist + model := &modelSystemObject{ + name: modelName, + } + + for _, o := range stmt.Options { + optName := parser.IdentName(o.Name) + + switch strings.ToLower(optName) { + case "modeltype": + lit, ok := o.OptionExpr.(*parser.StringLit) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", o.OptionExpr) + } + model.modelType = lit.Value + + case "labels": + lit, ok := o.OptionExpr.(*parser.SetLiteralExpr) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", o.OptionExpr) + } + model.labels = make([]string, len(lit.Members)) + for i, m := range lit.Members { + mlit, ok := m.(*parser.StringLit) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", m) + } + model.labels[i] = mlit.Value + } + + default: + return nil, sql3.NewErrInternalf("unexpected model option '%s'", optName) + } + } + + selOp, err := p.compileSelectStatement(stmt.ModelQuery, true) + if err != nil { + return nil, err + } + + // build a list of input columns for the model from the select query + schema := selOp.Schema() + model.inputColumns = make([]string, 0) + for _, p := range schema { + // if we have no column name, we have an error + if len(p.ColumnName) == 0 { + return nil, sql3.NewErrInternalf("query output columns used as inputs to models must be named") + } + // exclude any that are in the labels + isLabel := false + for _, l := range model.labels { + if strings.EqualFold(p.ColumnName, l) { + isLabel = true + break + } + } + if !isLabel { + model.inputColumns = append(model.inputColumns, p.ColumnName) + } + } + createModel := NewPlanOpCreateModel(p, model, selOp) + createModel.AddWarning("🦖 here there be dragons! CREATE MODEL statement is experimental.") + + query := NewPlanOpQuery(p, createModel, p.sql) + return query, nil +} + +func (p *ExecutionPlanner) analyzeCreateModelStatement(ctx context.Context, stmt *parser.CreateModelStatement) error { + // iterate the options + for _, opt := range stmt.Options { + optName := parser.IdentName(opt.Name) + if !isValidModelOption(optName) { + return sql3.NewErrInternalf("invalid model option '%s'", optName) + } + e, err := p.analyzeModelOptionExpr(ctx, optName, opt.OptionExpr, stmt) + if err != nil { + return err + } + opt.OptionExpr = e + } + + // analyze the select + _, err := p.analyzeSelectStatement(ctx, stmt.ModelQuery) + if err != nil { + return err + } + return nil +} + +func isValidModelOption(name string) bool { + switch strings.ToLower(name) { + case "modeltype": + return true + + case "labels": + return true + + default: + return false + } +} + +func (p *ExecutionPlanner) analyzeModelOptionExpr(ctx context.Context, optName string, expr parser.Expr, scope parser.Statement) (parser.Expr, error) { + if expr == nil { + return nil, nil + } + + e, err := p.analyzeExpression(ctx, expr, scope) + if err != nil { + return nil, err + } + + switch strings.ToLower(optName) { + case "modeltype": + + // model type needs to be a string literal + if !(e.IsLiteral() && typeIsString(e.DataType())) { + return nil, sql3.NewErrStringLiteral(e.Pos().Line, e.Pos().Column) + } + ty, ok := e.(*parser.StringLit) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", e) + } + + // these are the model types supported + switch strings.ToLower(ty.Value) { + case "linear_regresssion": + break + default: + return nil, sql3.NewErrInternalf("unexpected model tyoe '%s'", ty.Value) + } + return e, nil + + case "labels": + // labels needs to be a string array literal + // TODO (pok) revist 'set' literals (should be array literal; type checking could be robustified etc.) + if !e.IsLiteral() { + return nil, sql3.NewErrInternalf("string array literal expected") + } + ok, baseType := typeIsSet(e.DataType()) + if !ok { + return nil, sql3.NewErrInternalf("array expression expected") + } + if !typeIsString(baseType) { + return nil, sql3.NewErrInternalf("string array expected") + } + return e, nil + + default: + return nil, sql3.NewErrInternalf("unexpected option name '%s'", optName) + } +} diff --git a/sql3/planner/compiledropmodel.go b/sql3/planner/compiledropmodel.go new file mode 100644 index 000000000..e9d221230 --- /dev/null +++ b/sql3/planner/compiledropmodel.go @@ -0,0 +1,26 @@ +// Copyright 2023 Molecula Corp. All rights reserved. + +package planner + +import ( + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// compileDropModelStatement compiles a DROP MODEL statement into a PlanOperator. +func (p *ExecutionPlanner) compileDropModelStatement(stmt *parser.DropModelStatement) (_ types.PlanOperator, err error) { + modelName := parser.IdentName(stmt.Name) + v, err := p.getModelByName(modelName) + if err != nil { + return nil, err + } + if v == nil && !stmt.IfExists.IsValid() { + return nil, sql3.NewErrModelNotFound(0, 0, modelName) + } + + dropModel := NewPlanOpDropModel(p, stmt.IfExists.IsValid(), modelName) + dropModel.AddWarning("🦖 here there be dragons! DROP MODEL statement is experimental.") + + return NewPlanOpQuery(p, dropModel, p.sql), nil +} diff --git a/sql3/planner/compilepredict.go b/sql3/planner/compilepredict.go new file mode 100644 index 000000000..5219ddcea --- /dev/null +++ b/sql3/planner/compilepredict.go @@ -0,0 +1,50 @@ +// Copyright 2022 Molecula Corp. All rights reserved. + +package planner + +import ( + "context" + + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// compilePredictStatement compiles a parser.PredictStatement AST into a PlanOperator +func (p *ExecutionPlanner) compilePredictStatement(ctx context.Context, stmt *parser.PredictStatement) (types.PlanOperator, error) { + + // go get the model + modelName := parser.IdentName(stmt.ModelName) + + // does the model exist + obj, err := p.getModelByName(modelName) + if err != nil { + return nil, err + } + if obj == nil { + return nil, sql3.NewErrInternalf("model '%s' not found", modelName) + } + + selOp, err := p.compileSelectStatement(stmt.InputQuery, true) + if err != nil { + return nil, err + } + + predict := NewPlanOpPredict(p, obj, selOp) + predict.AddWarning("🦖 here there be dragons! PREDICT statement is experimental.") + + query := NewPlanOpQuery(p, predict, p.sql) + + return query, nil +} + +func (p *ExecutionPlanner) analyzePredictStatement(ctx context.Context, stmt *parser.PredictStatement) error { + + // analyze the select + _, err := p.analyzeSelectStatement(ctx, stmt.InputQuery) + if err != nil { + return err + } + + return nil +} diff --git a/sql3/planner/compileselect.go b/sql3/planner/compileselect.go index 2c986d3bb..c8e076c10 100644 --- a/sql3/planner/compileselect.go +++ b/sql3/planner/compileselect.go @@ -305,7 +305,7 @@ func (p *ExecutionPlanner) compileSelectStatement(stmt *parser.SelectStatement, } } - // insert the top operator if it exists + // insert the top operator if it exists, or limit - analyzer should have caught the case of both existing if stmt.Top.IsValid() { topExpr, err := p.compileExpr(stmt.TopExpr) if err != nil { @@ -313,6 +313,14 @@ func (p *ExecutionPlanner) compileSelectStatement(stmt *parser.SelectStatement, } compiledOp = NewPlanOpTop(topExpr, compiledOp) } + // handle limit + if stmt.Limit.IsValid() { + limitExpr, err := p.compileExpr(stmt.LimitExpr) + if err != nil { + return nil, err + } + compiledOp = NewPlanOpTop(limitExpr, compiledOp) + } // handle distinct if stmt.Distinct.IsValid() { @@ -613,6 +621,10 @@ func (p *ExecutionPlanner) analyzeSelectStatement(ctx context.Context, stmt *par } } + if stmt.TopExpr != nil && stmt.LimitExpr != nil { + return nil, sql3.NewErrErrTopLimitCannotCoexist(stmt.TopExpr.Pos().Line, stmt.TopExpr.Pos().Column) + } + expr, err := p.analyzeExpression(ctx, stmt.TopExpr, stmt) if err != nil { return nil, err @@ -624,6 +636,17 @@ func (p *ExecutionPlanner) analyzeSelectStatement(ctx context.Context, stmt *par stmt.TopExpr = expr } + expr, err = p.analyzeExpression(ctx, stmt.LimitExpr, stmt) + if err != nil { + return nil, err + } + if expr != nil { + if !(expr.IsLiteral() && typeIsInteger(expr.DataType())) { + return nil, sql3.NewErrIntegerLiteral(stmt.LimitExpr.Pos().Line, stmt.LimitExpr.Pos().Column) + } + stmt.LimitExpr = expr + } + expr, err = p.analyzeExpression(ctx, stmt.HavingExpr, stmt) if err != nil { return nil, err diff --git a/sql3/planner/executionplanner.go b/sql3/planner/executionplanner.go index 4f607173b..53cf69848 100644 --- a/sql3/planner/executionplanner.go +++ b/sql3/planner/executionplanner.go @@ -69,6 +69,10 @@ func (p *ExecutionPlanner) CompilePlan(ctx context.Context, stmt parser.Statemen rootOperator, err = p.compileSelectStatement(stmt, false) case *parser.ShowDatabasesStatement: rootOperator, err = p.compileShowDatabasesStatement(ctx, stmt) + case *parser.CopyStatement: + rootOperator, err = p.compileCopyStatement(stmt) + case *parser.PredictStatement: + rootOperator, err = p.compilePredictStatement(ctx, stmt) case *parser.ShowTablesStatement: rootOperator, err = p.compileShowTablesStatement(ctx, stmt) case *parser.ShowColumnsStatement: @@ -93,16 +97,23 @@ func (p *ExecutionPlanner) CompilePlan(ctx context.Context, stmt parser.Statemen rootOperator, err = p.compileDropTableStatement(ctx, stmt) case *parser.DropViewStatement: rootOperator, err = p.compileDropViewStatement(ctx, stmt) + case *parser.DropModelStatement: + rootOperator, err = p.compileDropModelStatement(stmt) case *parser.InsertStatement: rootOperator, err = p.compileInsertStatement(ctx, stmt) case *parser.BulkInsertStatement: rootOperator, err = p.compileBulkInsertStatement(ctx, stmt) case *parser.DeleteStatement: rootOperator, err = p.compileDeleteStatement(stmt) + case *parser.CreateModelStatement: + rootOperator, err = p.compileCreateModelStatement(stmt) + case *parser.CreateFunctionStatement: + rootOperator, err = p.compileCreateFunctionStatement(stmt) + default: return nil, sql3.NewErrInternalf("cannot plan statement: %T", stmt) } - // Optimize the plan. + // optimize the plan if err == nil { rootOperator, err = p.optimizePlan(ctx, rootOperator) } @@ -130,6 +141,10 @@ func (p *ExecutionPlanner) analyzePlan(ctx context.Context, stmt parser.Statemen return err case *parser.ShowDatabasesStatement: return nil + case *parser.CopyStatement: + return p.analyzeCopyStatement(ctx, stmt) + case *parser.PredictStatement: + return p.analyzePredictStatement(ctx, stmt) case *parser.ShowTablesStatement: return nil case *parser.ShowColumnsStatement: @@ -154,12 +169,19 @@ func (p *ExecutionPlanner) analyzePlan(ctx context.Context, stmt parser.Statemen return nil case *parser.DropViewStatement: return nil + case *parser.DropModelStatement: + return nil case *parser.InsertStatement: return p.analyzeInsertStatement(ctx, stmt) case *parser.BulkInsertStatement: return p.analyzeBulkInsertStatement(ctx, stmt) case *parser.DeleteStatement: return p.analyzeDeleteStatement(ctx, stmt) + case *parser.CreateModelStatement: + return p.analyzeCreateModelStatement(ctx, stmt) + case *parser.CreateFunctionStatement: + return p.analyzeCreateFunctionStatement(stmt) + default: return sql3.NewErrInternalf("cannot analyze statement: %T", stmt) } diff --git a/sql3/planner/expression.go b/sql3/planner/expression.go index 260e7052d..717f5486c 100644 --- a/sql3/planner/expression.go +++ b/sql3/planner/expression.go @@ -1514,16 +1514,18 @@ func (n *inOpPlanExpression) WithChildren(children ...types.PlanExpression) (typ // callPlanExpression is a function call type callPlanExpression struct { - name string - args []types.PlanExpression - dataType parser.ExprDataType + name string + args []types.PlanExpression + dataType parser.ExprDataType + udfReference *functionSystemObject } -func newCallPlanExpression(name string, args []types.PlanExpression, dataType parser.ExprDataType) *callPlanExpression { +func newCallPlanExpression(name string, args []types.PlanExpression, dataType parser.ExprDataType, udfReference *functionSystemObject) *callPlanExpression { return &callPlanExpression{ - name: name, - args: args, - dataType: dataType, + name: name, + args: args, + dataType: dataType, + udfReference: udfReference, } } @@ -1591,6 +1593,9 @@ func (n *callPlanExpression) Evaluate(currentRow []interface{}) (interface{}, er case "DATETIMEDIFF": return n.EvaluateDatetimeDiff(currentRow) default: + if n.udfReference != nil { + return n.evaluateUserDefinedFunction(currentRow) + } return nil, sql3.NewErrInternalf("unhandled function name '%s'", n.name) } } @@ -1632,7 +1637,7 @@ func (n *callPlanExpression) WithChildren(children ...types.PlanExpression) (typ if len(children) != len(n.args) { return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) } - return newCallPlanExpression(n.name, children, n.dataType), nil + return newCallPlanExpression(n.name, children, n.dataType, n.udfReference), nil } // aliasPlanExpression is a alias ref @@ -2049,30 +2054,30 @@ func (n *sysVariablePlanExpression) WithChildren(children ...types.PlanExpressio return n, nil } -// dateLiteralPlanExpression is a date literal -type dateLiteralPlanExpression struct { +// timestampLiteralPlanExpression is a date literal +type timestampLiteralPlanExpression struct { value time.Time } -func newDateLiteralPlanExpression(value time.Time) *dateLiteralPlanExpression { - return &dateLiteralPlanExpression{ +func newTimestampLiteralPlanExpression(value time.Time) *timestampLiteralPlanExpression { + return ×tampLiteralPlanExpression{ value: value, } } -func (n *dateLiteralPlanExpression) Evaluate(currentRow []interface{}) (interface{}, error) { +func (n *timestampLiteralPlanExpression) Evaluate(currentRow []interface{}) (interface{}, error) { return n.value, nil } -func (n *dateLiteralPlanExpression) Type() parser.ExprDataType { +func (n *timestampLiteralPlanExpression) Type() parser.ExprDataType { return parser.NewDataTypeTimestamp() } -func (n *dateLiteralPlanExpression) String() string { +func (n *timestampLiteralPlanExpression) String() string { return n.value.Format(time.RFC3339Nano) } -func (n *dateLiteralPlanExpression) Plan() map[string]interface{} { +func (n *timestampLiteralPlanExpression) Plan() map[string]interface{} { result := make(map[string]interface{}) result["_expr"] = fmt.Sprintf("%T", n) result["description"] = n.String() @@ -2081,11 +2086,11 @@ func (n *dateLiteralPlanExpression) Plan() map[string]interface{} { return result } -func (n *dateLiteralPlanExpression) Children() []types.PlanExpression { +func (n *timestampLiteralPlanExpression) Children() []types.PlanExpression { return []types.PlanExpression{} } -func (n *dateLiteralPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { +func (n *timestampLiteralPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { return n, nil } @@ -2660,7 +2665,7 @@ func (p *ExecutionPlanner) compileExpr(expr parser.Expr) (_ types.PlanExpression return newFloatLiteralPlanExpression(expr.Value), nil case *parser.DateLit: - return newDateLiteralPlanExpression(expr.Value), nil + return newTimestampLiteralPlanExpression(expr.Value), nil case *parser.SysVariable: return newSysVariablePlanExpression(expr.Name(), expr.Token), nil @@ -2892,6 +2897,14 @@ func (p *ExecutionPlanner) compileCallExpr(expr *parser.Call) (_ types.PlanExpre agg := newPercentilePlanExpression(args[0], args[1], expr.ResultDataType) return agg, nil + case "CORR": + agg := newCorrPlanExpression(args[0], args[1], expr.ResultDataType) + return agg, nil + + case "VAR": + agg := newVarPlanExpression(args[0], expr.ResultDataType) + return agg, nil + case "MIN": agg := newMinPlanExpression(args[0], expr.ResultDataType) return agg, nil @@ -2901,7 +2914,12 @@ func (p *ExecutionPlanner) compileCallExpr(expr *parser.Call) (_ types.PlanExpre return agg, nil default: - return newCallPlanExpression(parser.IdentName(expr.Name), args, expr.ResultDataType), nil + // could be a udf - try to look it up in functions + fn, err := p.getFunctionByName(strings.ToLower(callName)) + if err != nil { + return nil, err + } + return newCallPlanExpression(parser.IdentName(expr.Name), args, expr.ResultDataType, fn), nil } } diff --git a/sql3/planner/expression_it_test.go b/sql3/planner/expression_it_test.go index bf23baa34..d9f8d870b 100644 --- a/sql3/planner/expression_it_test.go +++ b/sql3/planner/expression_it_test.go @@ -40,7 +40,7 @@ func TestExpressions(t *testing.T) { iop := newInOpPlanExpression(newIntLiteralPlanExpression(10), parser.IN, newIntLiteralPlanExpression(20)) assert.Equal(t, iop.String(), "10 in (20)") - callop := newCallPlanExpression("foo", []types.PlanExpression{newIntLiteralPlanExpression(10)}, parser.NewDataTypeInt()) + callop := newCallPlanExpression("foo", []types.PlanExpression{newIntLiteralPlanExpression(10)}, parser.NewDataTypeInt(), nil) assert.Equal(t, callop.String(), "foo(10)") alop := newAliasPlanExpression("frobny", newIntLiteralPlanExpression(10)) @@ -65,7 +65,7 @@ func TestExpressions(t *testing.T) { assert.Equal(t, blop.String(), "false") tm, _ := time.ParseInLocation(time.RFC3339, "2012-11-01T22:08:41+00:00", time.UTC) - dlop := newDateLiteralPlanExpression(tm) + dlop := newTimestampLiteralPlanExpression(tm) assert.Equal(t, dlop.String(), "2012-11-01T22:08:41Z") slop := newStringLiteralPlanExpression("foo") diff --git a/sql3/planner/expressionagg.go b/sql3/planner/expressionagg.go index eab7f098d..b260a0c45 100644 --- a/sql3/planner/expressionagg.go +++ b/sql3/planner/expressionagg.go @@ -5,6 +5,7 @@ package planner import ( "context" "fmt" + "math" "github.com/featurebasedb/featurebase/v3/pql" "github.com/featurebasedb/featurebase/v3/sql3" @@ -943,6 +944,311 @@ func (n *percentilePlanExpression) WithChildren(children ...types.PlanExpression return newPercentilePlanExpression(children[0], children[1], n.returnDataType), nil } +// aggregator for CORR() +type aggregateCorr struct { + expr *corrPlanExpression + + n int64 + sum_X float64 + sum_Y float64 + sum_XY float64 + squareSum_X float64 + squareSum_Y float64 +} + +func NewAggCorrBuffer(child *corrPlanExpression) *aggregateCorr { + return &aggregateCorr{ + expr: child, + } +} + +func (m *aggregateCorr) Update(ctx context.Context, row types.Row) error { + v1, err := m.expr.arg1.Evaluate(row) + if err != nil { + return err + } + + v2, err := m.expr.arg2.Evaluate(row) + if err != nil { + return err + } + + // skip if nil + if v1 == nil || v2 == nil { + return nil + } + + var xVal float64 + var yVal float64 + + switch dataType := m.expr.arg1.Type().(type) { + case *parser.DataTypeDecimal: + thisVal, ok := v1.(pql.Decimal) + if !ok { + return sql3.NewErrInternalf("unexpected type conversion '%T'", v1) + } + + xVal = thisVal.Float64() + + case *parser.DataTypeInt: + thisVal, ok := v1.(int64) + if !ok { + return sql3.NewErrInternalf("unexpected type conversion '%T'", v1) + } + + xVal = float64(thisVal) + + default: + return sql3.NewErrInternalf("unhandled aggregate expression datatype '%T'", dataType) + } + + switch dataType := m.expr.arg2.Type().(type) { + case *parser.DataTypeDecimal: + thisVal, ok := v2.(pql.Decimal) + if !ok { + return sql3.NewErrInternalf("unexpected type conversion '%T'", v2) + } + + yVal = thisVal.Float64() + + case *parser.DataTypeInt: + thisVal, ok := v2.(int64) + if !ok { + return sql3.NewErrInternalf("unexpected type conversion '%T'", v2) + } + yVal = float64(thisVal) + + default: + return sql3.NewErrInternalf("unhandled aggregate expression datatype '%T'", dataType) + } + + m.sum_X = m.sum_X + xVal + m.sum_Y = m.sum_Y + yVal + + m.sum_XY = m.sum_XY + xVal*yVal + + m.squareSum_X = m.squareSum_X + xVal*xVal + m.squareSum_Y = m.squareSum_Y + yVal*yVal + m.n += 1 + + return nil +} + +func (m *aggregateCorr) Eval(ctx context.Context) (interface{}, error) { + corr := float64((float64(m.n)*m.sum_XY - m.sum_X*m.sum_Y)) / (math.Sqrt(float64((float64(m.n)*m.squareSum_X - m.sum_X*m.sum_X) * (float64(m.n)*m.squareSum_Y - m.sum_Y*m.sum_Y)))) + + d, err := pql.FromFloat64WithScale(corr, 6) + if err != nil { + return nil, err + } + return d, nil +} + +// corrPlanExpression handles CORR() - implement correlation coefficient +type corrPlanExpression struct { + arg1 types.PlanExpression + arg2 types.PlanExpression + returnDataType parser.ExprDataType +} + +var _ types.Aggregable = (*corrPlanExpression)(nil) + +func newCorrPlanExpression(arg1 types.PlanExpression, arg2 types.PlanExpression, returnDataType parser.ExprDataType) *corrPlanExpression { + return &corrPlanExpression{ + arg1: arg1, + arg2: arg2, + returnDataType: returnDataType, + } +} + +func (n *corrPlanExpression) Evaluate(currentRow []interface{}) (interface{}, error) { + return nil, sql3.NewErrInternalf("this should never be called") +} + +func (n *corrPlanExpression) NewBuffer() (types.AggregationBuffer, error) { + return NewAggCorrBuffer(n), nil +} + +func (n *corrPlanExpression) FirstChildExpr() types.PlanExpression { + return n.arg1 +} + +func (n *corrPlanExpression) Type() parser.ExprDataType { + return n.returnDataType +} + +func (n *corrPlanExpression) String() string { + return fmt.Sprintf("corr(%s, %s)", n.arg1.String(), n.arg2.String()) +} + +func (n *corrPlanExpression) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_expr"] = fmt.Sprintf("%T", n) + result["description"] = n.String() + result["dataType"] = n.Type().TypeDescription() + result["arg1"] = n.arg1.Plan() + result["arg2"] = n.arg2.Plan() + return result +} + +func (n *corrPlanExpression) Children() []types.PlanExpression { + return []types.PlanExpression{ + n.arg1, + n.arg2, + } +} + +func (n *corrPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { + if len(children) != 2 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newCorrPlanExpression(children[0], children[1], n.returnDataType), nil +} + +// aggregator for VAR() +type aggregateVar struct { + expr *varPlanExpression + + // to calculate mean + n int64 + sum float64 + + // we need to hang on to the values + // TODO(pok) - will need to spill these to disk for big result sets + values []float64 +} + +func NewAggVarBuffer(child *varPlanExpression) *aggregateVar { + return &aggregateVar{ + expr: child, + values: make([]float64, 0), + } +} + +func (m *aggregateVar) Update(ctx context.Context, row types.Row) error { + v, err := m.expr.arg.Evaluate(row) + if err != nil { + return err + } + + // skip if nil + if v == nil { + return nil + } + + var val float64 + + switch dataType := m.expr.arg.Type().(type) { + case *parser.DataTypeDecimal: + thisVal, ok := v.(pql.Decimal) + if !ok { + return sql3.NewErrInternalf("unexpected type conversion '%T'", v) + } + + val = thisVal.Float64() + + case *parser.DataTypeID: + thisVal, ok := v.(int64) + if !ok { + return sql3.NewErrInternalf("unexpected type conversion '%T'", v) + } + + val = float64(thisVal) + + case *parser.DataTypeInt: + thisVal, ok := v.(int64) + if !ok { + return sql3.NewErrInternalf("unexpected type conversion '%T'", v) + } + + val = float64(thisVal) + + default: + return sql3.NewErrInternalf("unhandled aggregate expression datatype '%T'", dataType) + } + + m.sum += val + m.n += 1 + m.values = append(m.values, val) + + return nil +} + +func (m *aggregateVar) Eval(ctx context.Context) (interface{}, error) { + + mean := m.sum / float64(m.n) + + var variance float64 + for _, v := range m.values { + variance += (v - mean) * (v - mean) + } + + variance = variance / float64(m.n) + + d, err := pql.FromFloat64WithScale(variance, 6) + if err != nil { + return nil, err + } + return d, nil +} + +// varPlanExpression handles VAR() - variance +type varPlanExpression struct { + arg types.PlanExpression + returnDataType parser.ExprDataType +} + +var _ types.Aggregable = (*varPlanExpression)(nil) + +func newVarPlanExpression(arg types.PlanExpression, returnDataType parser.ExprDataType) *varPlanExpression { + return &varPlanExpression{ + arg: arg, + returnDataType: returnDataType, + } +} + +func (n *varPlanExpression) Evaluate(currentRow []interface{}) (interface{}, error) { + return nil, sql3.NewErrInternalf("this should never be called") +} + +func (n *varPlanExpression) NewBuffer() (types.AggregationBuffer, error) { + return NewAggVarBuffer(n), nil +} + +func (n *varPlanExpression) FirstChildExpr() types.PlanExpression { + return n.arg +} + +func (n *varPlanExpression) Type() parser.ExprDataType { + return n.returnDataType +} + +func (n *varPlanExpression) String() string { + return fmt.Sprintf("var(%s)", n.arg.String()) +} + +func (n *varPlanExpression) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_expr"] = fmt.Sprintf("%T", n) + result["description"] = n.String() + result["dataType"] = n.Type().TypeDescription() + result["arg"] = n.arg.Plan() + return result +} + +func (n *varPlanExpression) Children() []types.PlanExpression { + return []types.PlanExpression{ + n.arg, + } +} + +func (n *varPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newVarPlanExpression(children[0], n.returnDataType), nil +} + // aggregator for LAST() type aggregateLast struct { val interface{} diff --git a/sql3/planner/expressionanalyzer.go b/sql3/planner/expressionanalyzer.go index 1adf414c4..76ce2068a 100644 --- a/sql3/planner/expressionanalyzer.go +++ b/sql3/planner/expressionanalyzer.go @@ -632,7 +632,7 @@ func (p *ExecutionPlanner) analyzeBinaryExpression(ctx context.Context, expr *pa if ok { //we have a select in the expression list so make sure it is the only thing in the expression list if len(lst.Exprs) > 1 { - return nil, sql3.NewErrInternalf("expresion list should only contain one select statement") + return nil, sql3.NewErrInternalf("expression list should only contain one select statement") } //make sure select only returns one column if len(sel.Columns) > 1 { diff --git a/sql3/planner/expressionanalyzercall.go b/sql3/planner/expressionanalyzercall.go index 597fbd1ba..b75c2c8fa 100644 --- a/sql3/planner/expressionanalyzercall.go +++ b/sql3/planner/expressionanalyzercall.go @@ -124,6 +124,71 @@ func (p *ExecutionPlanner) analyzeCallExpression(ctx context.Context, call *pars //return the data type of the referenced column call.ResultDataType = ref.DataType() + case "CORR": + // can't do this on a * + if call.Star.IsValid() && len(call.Args) == 0 { + return nil, sql3.NewErrExpectedColumnReference(call.Star.Line, call.Star.Column) + } + + if len(call.Args) != 2 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 2, len(call.Args)) + } + + // if it is a ref, we shouldn't do a corr on the _id + arg1 := call.Args[0] + + ref, ok := arg1.(*parser.QualifiedRef) + if ok && strings.EqualFold(ref.Column.Name, string(dax.PrimaryKeyFieldName)) { + return nil, sql3.NewErrIdColumnNotValidForAggregateFunction(call.Args[0].Pos().Line, call.Args[0].Pos().Column, call.Name.Name) + } + + // make sure the ref is the right type + if !(typeIsInteger(arg1.DataType()) || typeIsDecimal(arg1.DataType()) || typeIsTimestamp(arg1.DataType())) { + return nil, sql3.NewErrIntOrDecimalOrTimestampExpressionExpected(arg1.Pos().Line, arg1.Pos().Column) + } + + // if it is a ref, we shouldn't do a corr on the _id + arg2 := call.Args[1] + + ref, ok = arg2.(*parser.QualifiedRef) + if ok && strings.EqualFold(ref.Column.Name, string(dax.PrimaryKeyFieldName)) { + return nil, sql3.NewErrIdColumnNotValidForAggregateFunction(call.Args[1].Pos().Line, call.Args[1].Pos().Column, call.Name.Name) + } + + // make sure the ref is the right type + if !(typeIsInteger(arg2.DataType()) || typeIsDecimal(arg2.DataType()) || typeIsTimestamp(arg2.DataType())) { + return nil, sql3.NewErrIntOrDecimalOrTimestampExpressionExpected(arg2.Pos().Line, arg2.Pos().Column) + } + + // return the data type of the referenced column + call.ResultDataType = parser.NewDataTypeDecimal(6) + + case "VAR": + // can't do this on a * + if call.Star.IsValid() && len(call.Args) == 0 { + return nil, sql3.NewErrExpectedColumnReference(call.Star.Line, call.Star.Column) + } + + if len(call.Args) != 1 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) + } + + // first arg should be a qualified ref + arg1 := call.Args[0] + + ref, ok := arg1.(*parser.QualifiedRef) + if ok && strings.EqualFold(ref.Column.Name, string(dax.PrimaryKeyFieldName)) { + return nil, sql3.NewErrIdColumnNotValidForAggregateFunction(call.Args[0].Pos().Line, call.Args[0].Pos().Column, call.Name.Name) + } + + // make sure the ref is the right type + if !(typeIsInteger(arg1.DataType()) || typeIsDecimal(arg1.DataType()) || typeIsTimestamp(arg1.DataType())) { + return nil, sql3.NewErrIntOrDecimalOrTimestampExpressionExpected(arg1.Pos().Line, arg1.Pos().Column) + } + + // return the data type of the referenced column + call.ResultDataType = parser.NewDataTypeDecimal(6) + case "MIN", "MAX": // can't do an min/max on a * if call.Star.IsValid() && len(call.Args) == 0 { @@ -270,6 +335,15 @@ func (p *ExecutionPlanner) analyzeCallExpression(ctx context.Context, call *pars case "DATETIMEDIFF": return p.analyzeFunctionDateTimeDiff(call, scope) default: + // could be a udf - try to look it up in functions + fn, err := p.getFunctionByName(call.Name.Name) + if err != nil { + return nil, err + } + if fn != nil { + return p.analyzeUserDefinedFunction(call, scope, fn) + } + return nil, sql3.NewErrCallUnknownFunction(call.Name.NamePos.Line, call.Name.NamePos.Column, call.Name.Name) } return call, nil diff --git a/sql3/planner/expressionpql.go b/sql3/planner/expressionpql.go index c6f2022b7..48526a841 100644 --- a/sql3/planner/expressionpql.go +++ b/sql3/planner/expressionpql.go @@ -604,7 +604,7 @@ func planExprToValue(expr types.PlanExpression) (interface{}, error) { return expr.value, nil case *stringLiteralPlanExpression: return expr.value, nil - case *dateLiteralPlanExpression: + case *timestampLiteralPlanExpression: return expr.value, nil case *boolLiteralPlanExpression: return expr.value, nil diff --git a/sql3/planner/opaltertable.go b/sql3/planner/opaltertable.go index 6da8db0e5..2b6495843 100644 --- a/sql3/planner/opaltertable.go +++ b/sql3/planner/opaltertable.go @@ -94,11 +94,11 @@ func (i *alterTableRowIter) Next(ctx context.Context) (types.Row, error) { fos := i.columnDef.fos fld, err := pilosa.FieldFromFieldOptions(fname, fos...) - // all newly created fields unconditionally have TrackExistence turned on. - fld.Options.TrackExistence = true if err != nil { return nil, err } + // all newly created fields unconditionally have TrackExistence turned on. + fld.Options.TrackExistence = true if err := i.planner.schemaAPI.CreateField(ctx, tname, fld); err != nil { return nil, err diff --git a/sql3/planner/opbulkinsert.go b/sql3/planner/opbulkinsert.go index a8d309cb8..2586ad34b 100644 --- a/sql3/planner/opbulkinsert.go +++ b/sql3/planner/opbulkinsert.go @@ -867,7 +867,7 @@ func processColumnValue(rawValue interface{}, targetType parser.ExprDataType) (t if !ok { return nil, sql3.NewErrInternalf("unable to convert '%s", rawValue) } - return newDateLiteralPlanExpression(tval), nil + return newTimestampLiteralPlanExpression(tval), nil case *parser.DataTypeString: sval, ok := rawValue.(string) diff --git a/sql3/planner/opcopy.go b/sql3/planner/opcopy.go new file mode 100644 index 000000000..2ede81ff0 --- /dev/null +++ b/sql3/planner/opcopy.go @@ -0,0 +1,515 @@ +// Copyright 2022 Molecula Corp. All rights reserved. + +package planner + +import ( + "bytes" + "context" + "fmt" + "io" + "net/http" + "strconv" + "strings" + "time" + + pilosa "github.com/featurebasedb/featurebase/v3" + "github.com/featurebasedb/featurebase/v3/pql" + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// PlanOpCopy is a copy operator +type PlanOpCopy struct { + planner *ExecutionPlanner + targetTable string + url string + apiKey string + ddl string + ChildOp types.PlanOperator + + warnings []string +} + +func NewPlanOpCopy(planner *ExecutionPlanner, targetName string, url string, apiKey string, ddl string, child types.PlanOperator) *PlanOpCopy { + return &PlanOpCopy{ + planner: planner, + targetTable: targetName, + url: url, + apiKey: apiKey, + ddl: ddl, + ChildOp: child, + warnings: make([]string, 0), + } +} + +func (p *PlanOpCopy) Schema() types.Schema { + return types.Schema{} +} + +func (p *PlanOpCopy) Iterator(ctx context.Context, row types.Row) (types.RowIterator, error) { + child, err := p.ChildOp.Iterator(ctx, row) + if err != nil { + return nil, err + } + if p.url != "" { + return newRemoteCopyIterator(p.planner, p.targetTable, p.url, p.apiKey, p.ddl, p.ChildOp.Schema(), child), nil + } + return newCopyIterator(p.planner, p.targetTable, p.ddl, p.ChildOp.Schema(), child), nil +} + +func (p *PlanOpCopy) WithChildren(children ...types.PlanOperator) (types.PlanOperator, error) { + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return NewPlanOpCopy(p.planner, p.targetTable, p.url, p.apiKey, p.ddl, children[0]), nil +} + +func (p *PlanOpCopy) Children() []types.PlanOperator { + return []types.PlanOperator{ + p.ChildOp, + } +} + +func (p *PlanOpCopy) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_op"] = fmt.Sprintf("%T", p) + result["_schema"] = p.Schema().Plan() + result["child"] = p.ChildOp.Plan() + result["child"] = p.ChildOp.Plan() + return result +} + +func (p *PlanOpCopy) String() string { + return "" +} + +func (p *PlanOpCopy) AddWarning(warning string) { + p.warnings = append(p.warnings, warning) +} + +func (p *PlanOpCopy) Warnings() []string { + return p.warnings +} + +func (p *PlanOpCopy) Expressions() []types.PlanExpression { + return []types.PlanExpression{} +} + +func (p *PlanOpCopy) WithUpdatedExpressions(exprs ...types.PlanExpression) (types.PlanOperator, error) { + if len(exprs) > 0 { + return nil, sql3.NewErrInternalf("unexpected number of exprs '%d'", len(exprs)) + } + return p, nil +} + +type copyIterator struct { + planner *ExecutionPlanner + targetTableName string + copySchema types.Schema + ddl string + child types.RowIterator + hasStarted *struct{} +} + +func newCopyIterator(planner *ExecutionPlanner, targetTableName string, ddl string, copySchema types.Schema, childIter types.RowIterator) *copyIterator { + return ©Iterator{ + planner: planner, + targetTableName: targetTableName, + ddl: ddl, + copySchema: copySchema, + child: childIter, + } +} + +func (i *copyIterator) Next(ctx context.Context) (types.Row, error) { + if i.hasStarted == nil { + // parse and execute the ddl to create the table + ast, err := parser.NewParser(strings.NewReader(i.ddl)).ParseStatement() + if err != nil { + return nil, err + } + ct, ok := ast.(*parser.CreateTableStatement) + if !ok { + return nil, sql3.NewErrInternalf("unexpected ast type") + } + // analyze + err = i.planner.analyzeCreateTableStatement(ct) + if err != nil { + return nil, err + } + ctOp, err := i.planner.compileCreateTableStatement(ctx, ct) + if err != nil { + return nil, err + } + ctIter, err := ctOp.Iterator(context.Background(), nil) + if err != nil { + return nil, err + } + _, err = ctIter.Next(ctx) + if err != nil && err != types.ErrNoMoreRows { + return nil, err + } + + targetColumns := make([]*qualifiedRefPlanExpression, 0) + + for _, s := range i.copySchema { + targetColumns = append(targetColumns, newQualifiedRefPlanExpression(i.targetTableName, s.ColumnName, 0, s.Type)) + } + + // build an insert iterator for the target table + insertIter := &insertRowIter{ + planner: i.planner, + tableName: i.targetTableName, + targetColumns: targetColumns, + } + + batchCount := 0 + insertBatch := make([][]types.PlanExpression, 0) + + for { + // get a source row + row, err := i.child.Next(ctx) + if err != nil { + if err == types.ErrNoMoreRows { + break + } + return nil, err + } + + // add it to target batch + + irow := make([]types.PlanExpression, len(row)) + for i, s := range i.copySchema { + switch ty := s.Type.(type) { + case *parser.DataTypeID, *parser.DataTypeInt: + val, ok := row[i].(int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + irow[i] = newIntLiteralPlanExpression(val) + + case *parser.DataTypeDecimal: + val, ok := row[i].(pql.Decimal) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + irow[i] = newFloatLiteralPlanExpression(val.String()) + + case *parser.DataTypeString: + val, ok := row[i].(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + irow[i] = newStringLiteralPlanExpression(val) + + case *parser.DataTypeBool: + val, ok := row[i].(bool) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + irow[i] = newBoolLiteralPlanExpression(val) + + case *parser.DataTypeTimestamp: + val, ok := row[i].(time.Time) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + irow[i] = newTimestampLiteralPlanExpression(val) + + case *parser.DataTypeStringSet: + val, ok := row[i].([]string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + + members := make([]types.PlanExpression, 0) + for _, m := range val { + members = append(members, newStringLiteralPlanExpression(m)) + } + irow[i] = newExprSetLiteralPlanExpression(members, parser.NewDataTypeStringSet()) + + case *parser.DataTypeIDSet: + val, ok := row[i].([]int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + + members := make([]types.PlanExpression, 0) + for _, m := range val { + members = append(members, newIntLiteralPlanExpression(m)) + } + irow[i] = newExprSetLiteralPlanExpression(members, parser.NewDataTypeIDSet()) + + default: + return nil, sql3.NewErrInternalf("unhandled type '%T'", ty) + } + } + insertBatch = append(insertBatch, irow) + + // inc batch count + batchCount += 1 + if batchCount > 1000 { + // do the insert + insertIter.insertValues = insertBatch + _, err = insertIter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return nil, err + } + // reset + batchCount = 0 + insertBatch = make([][]types.PlanExpression, 0) + } + } + if len(insertBatch) > 0 { + // do the insert + insertIter.insertValues = insertBatch + _, err = insertIter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return nil, err + } + } + + i.hasStarted = &struct{}{} + } + return nil, types.ErrNoMoreRows +} + +type remoteCopyIterator struct { + planner *ExecutionPlanner + targetTableName string + copySchema types.Schema + url string + apiKey string + ddl string + child types.RowIterator + hasStarted *struct{} +} + +func newRemoteCopyIterator(planner *ExecutionPlanner, targetTableName string, url string, apiKey string, ddl string, copySchema types.Schema, childIter types.RowIterator) *remoteCopyIterator { + return &remoteCopyIterator{ + planner: planner, + targetTableName: targetTableName, + url: url, + apiKey: apiKey, + ddl: ddl, + copySchema: copySchema, + child: childIter, + } +} + +func (i *remoteCopyIterator) remoteExec(ctx context.Context, sql string) (*pilosa.WireQueryResponse, error) { + // Create HTTP request. + req, err := http.NewRequest("POST", i.url, strings.NewReader(sql)) + if err != nil { + return nil, sql3.NewErrInternalf("error executing remotely: %s", err.Error()) + } + + req.Header.Set("Content-Length", strconv.Itoa(len(sql))) + req.Header.Set("Content-Type", "text/plain") + req.Header.Set("Accept", "application/json") + req.Header.Set("User-Agent", "pilosa/"+i.planner.systemAPI.Version()) + if len(i.apiKey) > 0 { + req.Header.Set("X-API-Key", i.apiKey) + } + + // Execute request against the host. + resp, err := http.DefaultClient.Do(req.WithContext(ctx)) + if err != nil { + return nil, sql3.NewErrInternalf("error executing remotely: %s", err.Error()) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, sql3.NewErrInternalf("error executing remotely: %s", err.Error()) + } + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + if resp.StatusCode == 401 { + return nil, sql3.NewErrRemoteUnauthorized(0, 0, i.url) + } + // we have an error + return nil, sql3.NewErrInternalf("error executing remotely: %d, %s", resp.StatusCode, string(body)) + } + + sqlResponse := &pilosa.WireQueryResponse{} + err = sqlResponse.UnmarshalJSONTyped([]byte(body), true) + if err != nil { + return nil, sql3.NewErrInternalf("error executing remotely: %s", err.Error()) + } + + if len(sqlResponse.Error) > 0 { + return nil, sql3.NewErrInternalf("error executing remotely: %s", sqlResponse.Error) + } + + return sqlResponse, nil +} + +func (i *remoteCopyIterator) Next(ctx context.Context) (types.Row, error) { + if i.hasStarted == nil { + // execute the ddl to create the table + _, err := i.remoteExec(ctx, i.ddl) + if err != nil { + return nil, err + } + + // build bulk insert statement + var buf bytes.Buffer + buf.WriteString("bulk insert into ") + fmt.Fprintf(&buf, "%s", i.targetTableName) + buf.WriteString(" (") + + for i, s := range i.copySchema { + if i > 0 { + buf.WriteString(", ") + } + fmt.Fprintf(&buf, "%s", s.ColumnName) + } + buf.WriteString(") map (") + for i, s := range i.copySchema { + if i > 0 { + buf.WriteString(", ") + } + fmt.Fprintf(&buf, "'$._%d' %s", i, s.Type.TypeDescription()) + } + buf.WriteString(") from x'") + header := buf.String() + + batchCount := 0 + var batchBuf bytes.Buffer + + for { + // get a source row + row, err := i.child.Next(ctx) + if err != nil { + if err == types.ErrNoMoreRows { + break + } + return nil, err + } + + // add it to target batch + var rowBuf bytes.Buffer + rowBuf.WriteString("{") + for i, s := range i.copySchema { + if i > 0 { + rowBuf.WriteString(",") + } + fmt.Fprintf(&rowBuf, `"_%d":`, i) + + if row[i] == nil { + rowBuf.WriteString("null") + continue + } + + switch ty := s.Type.(type) { + case *parser.DataTypeID, *parser.DataTypeInt: + val, ok := row[i].(int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + fmt.Fprintf(&rowBuf, "%d", val) + + case *parser.DataTypeString: + val, ok := row[i].(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + // escape single quotes + val = strings.ReplaceAll(val, `'`, `''`) + // and double quotes + val = strings.ReplaceAll(val, `"`, `\"`) + // and line feeds + if strings.Contains(val, "\n") { + val = strings.ReplaceAll(val, "\n", "\\n") + } + fmt.Fprintf(&rowBuf, `"%s"`, val) + + case *parser.DataTypeBool: + val, ok := row[i].(bool) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + fmt.Fprintf(&rowBuf, "%v", val) + + case *parser.DataTypeTimestamp: + val, ok := row[i].(time.Time) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + fmt.Fprintf(&rowBuf, `"%s"`, val.Format(time.RFC3339Nano)) + + case *parser.DataTypeStringSet: + val, ok := row[i].([]string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + rowBuf.WriteString("[") + for j, s := range val { + if j > 0 { + rowBuf.WriteString(",") + } + fmt.Fprintf(&rowBuf, `"%s"`, s) + } + rowBuf.WriteString("]") + + case *parser.DataTypeIDSet: + val, ok := row[i].([]int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type '%T'", row[i]) + } + rowBuf.WriteString("[") + for j, s := range val { + if j > 0 { + rowBuf.WriteString(",") + } + fmt.Fprintf(&rowBuf, `%d`, s) + } + rowBuf.WriteString("]") + + default: + return nil, sql3.NewErrInternalf("unhandled type '%T'", ty) + } + } + rowBuf.WriteString("}\n") + batchBuf.Write(rowBuf.Bytes()) + + // inc batch count + batchCount += 1 + if batchCount > 10000 { + // do the insert + + var reqBuf bytes.Buffer + reqBuf.WriteString(header) + reqBuf.Write(batchBuf.Bytes()) + reqBuf.WriteString("' with batchsize 10000 input 'STREAM' format 'NDJSON'") + + _, err := i.remoteExec(ctx, reqBuf.String()) + if err != nil { + return nil, err + } + + // reset + batchCount = 0 + batchBuf.Reset() + } + } + if batchCount > 0 { + // do the insert + + var reqBuf bytes.Buffer + reqBuf.WriteString(header) + reqBuf.Write(batchBuf.Bytes()) + reqBuf.WriteString("' with batchsize 10000 input 'STREAM' format 'NDJSON'") + + _, err := i.remoteExec(ctx, reqBuf.String()) + if err != nil { + return nil, err + } + } + + i.hasStarted = &struct{}{} + } + return nil, types.ErrNoMoreRows +} diff --git a/sql3/planner/opcreatefunction.go b/sql3/planner/opcreatefunction.go new file mode 100644 index 000000000..0c84c88aa --- /dev/null +++ b/sql3/planner/opcreatefunction.go @@ -0,0 +1,104 @@ +// Copyright 2023 Molecula Corp. All rights reserved. + +package planner + +import ( + "context" + "fmt" + + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// PlanOpCreateFunction implements the CREATE FUNCTION operator +type PlanOpCreateFunction struct { + planner *ExecutionPlanner + function *functionSystemObject + ifNotExists bool + warnings []string +} + +func NewPlanOpCreateFunction(planner *ExecutionPlanner, ifNotExists bool, function *functionSystemObject) *PlanOpCreateFunction { + return &PlanOpCreateFunction{ + planner: planner, + function: function, + ifNotExists: ifNotExists, + warnings: make([]string, 0), + } +} + +func (p *PlanOpCreateFunction) Schema() types.Schema { + return types.Schema{} +} + +func (p *PlanOpCreateFunction) Iterator(ctx context.Context, row types.Row) (types.RowIterator, error) { + return newCreateFunctionIter(p.planner, p.ifNotExists, p.function), nil +} + +func (p *PlanOpCreateFunction) Children() []types.PlanOperator { + return []types.PlanOperator{} +} + +func (p *PlanOpCreateFunction) WithChildren(children ...types.PlanOperator) (types.PlanOperator, error) { + if len(children) != 0 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return NewPlanOpCreateFunction(p.planner, p.ifNotExists, p.function), nil +} + +func (p *PlanOpCreateFunction) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_op"] = fmt.Sprintf("%T", p) + result["_schema"] = p.Schema().Plan() + result["model"] = p.function.name + return result +} + +func (p *PlanOpCreateFunction) String() string { + return "" +} + +func (p *PlanOpCreateFunction) AddWarning(warning string) { + p.warnings = append(p.warnings, warning) +} + +func (p *PlanOpCreateFunction) Warnings() []string { + var w []string + w = append(w, p.warnings...) + return w +} + +type createFunctionIter struct { + planner *ExecutionPlanner + function *functionSystemObject + ifNotExists bool +} + +func newCreateFunctionIter(planner *ExecutionPlanner, ifNotExists bool, function *functionSystemObject) *createFunctionIter { + return &createFunctionIter{ + planner: planner, + function: function, + ifNotExists: ifNotExists, + } +} + +func (i *createFunctionIter) Next(ctx context.Context) (types.Row, error) { + // now check in the functions table to see if it is exists + v, err := i.planner.getFunctionByName(i.function.name) + if err != nil { + return nil, err + } + if v != nil { + if i.ifNotExists { + return nil, types.ErrNoMoreRows + } + return nil, sql3.NewErrViewExists(0, 0, i.function.name) + } + + // now store the view into fb_functions + err = i.planner.insertFunction(i.function) + if err != nil { + return nil, err + } + return nil, types.ErrNoMoreRows +} diff --git a/sql3/planner/opcreatemodel.go b/sql3/planner/opcreatemodel.go new file mode 100644 index 000000000..39136e0ca --- /dev/null +++ b/sql3/planner/opcreatemodel.go @@ -0,0 +1,279 @@ +// Copyright 2022 Molecula Corp. All rights reserved. + +package planner + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "github.com/featurebasedb/featurebase/v3/pql" + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" + uuid "github.com/satori/go.uuid" +) + +// PlanOpCreateModel implements the CREATE MODEL operator +type PlanOpCreateModel struct { + ChildOp types.PlanOperator + planner *ExecutionPlanner + model *modelSystemObject + warnings []string +} + +func NewPlanOpCreateModel(planner *ExecutionPlanner, model *modelSystemObject, child types.PlanOperator) *PlanOpCreateModel { + return &PlanOpCreateModel{ + ChildOp: child, + planner: planner, + model: model, + warnings: make([]string, 0), + } +} + +func (p *PlanOpCreateModel) Schema() types.Schema { + return types.Schema{} +} + +func (p *PlanOpCreateModel) Iterator(ctx context.Context, row types.Row) (types.RowIterator, error) { + // get the query iterator + iter, err := p.ChildOp.Iterator(ctx, row) + if err != nil { + return nil, err + } + + switch strings.ToLower(p.model.modelType) { + case "linear_regresssion": + return newCreateModelIter(p.planner, p.model, newLinearRegressionModelIter(p.planner, p.model, p.ChildOp.Schema(), iter)), nil + + default: + return nil, sql3.NewErrInternalf("unexpected model tyoe '%s'", p.model.modelType) + } +} + +func (p *PlanOpCreateModel) Children() []types.PlanOperator { + return []types.PlanOperator{ + p.ChildOp, + } +} + +func (p *PlanOpCreateModel) WithChildren(children ...types.PlanOperator) (types.PlanOperator, error) { + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return NewPlanOpCreateModel(p.planner, p.model, children[0]), nil +} + +func (p *PlanOpCreateModel) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_op"] = fmt.Sprintf("%T", p) + sc := make([]string, 0) + for _, e := range p.Schema() { + sc = append(sc, fmt.Sprintf("'%s', '%s', '%s'", e.ColumnName, e.RelationName, e.Type.TypeDescription())) + } + result["_schema"] = sc + result["model"] = p.model.name // TODO(pok) - add a Plan() method here (or some such) + result["child"] = p.ChildOp.Plan() + return result +} + +func (p *PlanOpCreateModel) String() string { + return "" +} + +func (p *PlanOpCreateModel) AddWarning(warning string) { + p.warnings = append(p.warnings, warning) +} + +func (p *PlanOpCreateModel) Warnings() []string { + var w []string + w = append(w, p.warnings...) + w = append(w, p.ChildOp.Warnings()...) + return w +} + +type createModelIter struct { + child types.RowIterator + planner *ExecutionPlanner + model *modelSystemObject + hasStarted *struct{} +} + +func newCreateModelIter(planner *ExecutionPlanner, model *modelSystemObject, child types.RowIterator) *createModelIter { + return &createModelIter{ + planner: planner, + model: model, + child: child, + } +} + +func (i *createModelIter) Next(ctx context.Context) (types.Row, error) { + if i.hasStarted == nil { + // store the model into fb_models and set the model status to 'training' + i.model.status = "TRAINING" + err := i.planner.insertModel(i.model) + if err != nil { + return nil, err + } + + // do the actual training + _, err = i.child.Next(ctx) + if err != nil && err != types.ErrNoMoreRows { + return nil, err + } + + // update the model to ready + i.model.status = "READY" + err = i.planner.updateModel(i.model) + if err != nil { + return nil, err + } + i.hasStarted = &struct{}{} + } + return nil, types.ErrNoMoreRows +} + +type linearRegressionModelIter struct { + child types.RowIterator + planner *ExecutionPlanner + model *modelSystemObject + childSchema types.Schema + hasStarted *struct{} +} + +func newLinearRegressionModelIter(planner *ExecutionPlanner, model *modelSystemObject, childSchema types.Schema, child types.RowIterator) *linearRegressionModelIter { + return &linearRegressionModelIter{ + planner: planner, + model: model, + childSchema: childSchema, + child: child, + } +} + +func (i *linearRegressionModelIter) Next(ctx context.Context) (types.Row, error) { + if i.hasStarted == nil { + // this is linear regression now, so we actually 'train' when we predict (later) + // for now we just store the values from the query in fb_model_data + + // delete anything from fb_model_data for this model + err := i.planner.ensureModelDataSystemTableExists() + if err != nil { + return nil, err + } + diter := &filteredDeleteRowIter{ + planner: i.planner, + tableName: "fb_model_data", + filter: newBinOpPlanExpression( + newQualifiedRefPlanExpression("fb_model_data", "model_id", 0, parser.NewDataTypeString()), + parser.EQ, + newStringLiteralPlanExpression(i.model.name), + parser.NewDataTypeBool(), + ), + } + _, err = diter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return nil, err + } + + iter := &insertRowIter{ + planner: i.planner, + tableName: "fb_model_data", + targetColumns: []*qualifiedRefPlanExpression{ + newQualifiedRefPlanExpression("fb_model_data", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_model_data", "model_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_model_data", "data", 0, parser.NewDataTypeString()), + }, + insertValues: [][]types.PlanExpression{}, + } + + trainingRefs := make([]*qualifiedRefPlanExpression, 0) + + // make sure label column exists and is type compatible with float + labelColumn := i.model.labels[0] + found := false + for i, s := range i.childSchema { + if strings.EqualFold(labelColumn, s.ColumnName) { + if !typesAreAssignmentCompatible(parser.NewDataTypeDecimal(4), s.Type) { + return nil, sql3.NewErrInternalf("types not assignment compatible") + } + trainingRefs = append(trainingRefs, newQualifiedRefPlanExpression("", labelColumn, i, s.Type)) + found = true + break + } + } + if !found { + return nil, sql3.NewErrInternalf("label column found found") + } + + // make sure input columns exists and are type compatible with float + for _, ic := range i.model.inputColumns { + found := false + for i, s := range i.childSchema { + if strings.EqualFold(ic, s.ColumnName) { + if !typesAreAssignmentCompatible(parser.NewDataTypeDecimal(4), s.Type) { + return nil, sql3.NewErrInternalf("types not assignment compatible") + } + trainingRefs = append(trainingRefs, newQualifiedRefPlanExpression("", ic, i, s.Type)) + found = true + break + } + } + if !found { + return nil, sql3.NewErrInternalf("input column found found") + } + } + + // go run the query and iterate + for { + row, err := i.child.Next(ctx) + if err != nil { + if err == types.ErrNoMoreRows { + break + } + return nil, err + } + + fdata := make([]float64, 0) + + for _, ref := range trainingRefs { + val, err := ref.Evaluate(row) + if err != nil { + return nil, err + } + cval, err := coerceValue(ref.dataType, parser.NewDataTypeDecimal(4), val, parser.Pos{Line: 0, Column: 0}) + if err != nil { + return nil, err + } + dval := cval.(pql.Decimal) + fdata = append(fdata, dval.Float64()) + } + + data, err := json.Marshal(fdata) + if err != nil { + return nil, err + } + + rowID, err := uuid.NewV4() + if err != nil { + return nil, err + } + tuple := []types.PlanExpression{ + newStringLiteralPlanExpression(rowID.String()), + newStringLiteralPlanExpression(i.model.name), + newStringLiteralPlanExpression(string(data)), + } + iter.insertValues = append(iter.insertValues, tuple) + + fmt.Printf("%v", row) + + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return nil, err + } + + i.hasStarted = &struct{}{} + } + return nil, types.ErrNoMoreRows +} diff --git a/sql3/planner/opcreatetable.go b/sql3/planner/opcreatetable.go index 25de51bbf..5b143c9b6 100644 --- a/sql3/planner/opcreatetable.go +++ b/sql3/planner/opcreatetable.go @@ -112,11 +112,11 @@ func (i *createTableRowIter) Next(ctx context.Context) (types.Row, error) { for _, f := range i.columns { fld, err := pilosa.FieldFromFieldOptions(dax.FieldName(f.name), f.fos...) - // We unconditionally turn on TrackExistence for all newly-created fields. - fld.Options.TrackExistence = true if err != nil { return nil, errors.Wrapf(err, "creating field from field options: %s", f.name) } + // We unconditionally turn on TrackExistence for all newly-created fields. + fld.Options.TrackExistence = true fields = append(fields, fld) } diff --git a/sql3/planner/opdropmodel.go b/sql3/planner/opdropmodel.go new file mode 100644 index 000000000..86cc3ee7b --- /dev/null +++ b/sql3/planner/opdropmodel.go @@ -0,0 +1,102 @@ +// Copyright 2023 Molecula Corp. All rights reserved. + +package planner + +import ( + "context" + "fmt" + + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" +) + +// PlanOpDropModel plan operator to drop a view. +type PlanOpDropModel struct { + planner *ExecutionPlanner + modelName string + ifExists bool + warnings []string +} + +func NewPlanOpDropModel(p *ExecutionPlanner, ifExists bool, modelName string) *PlanOpDropModel { + return &PlanOpDropModel{ + planner: p, + modelName: modelName, + ifExists: ifExists, + warnings: make([]string, 0), + } +} + +func (p *PlanOpDropModel) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_op"] = fmt.Sprintf("%T", p) + result["modelName"] = p.modelName + result["isExists"] = p.ifExists + return result +} + +func (p *PlanOpDropModel) String() string { + return "" +} + +func (p *PlanOpDropModel) AddWarning(warning string) { + p.warnings = append(p.warnings, warning) +} + +func (p *PlanOpDropModel) Warnings() []string { + return p.warnings +} + +func (p *PlanOpDropModel) Schema() types.Schema { + return types.Schema{} +} + +func (p *PlanOpDropModel) Children() []types.PlanOperator { + return []types.PlanOperator{} +} + +func (p *PlanOpDropModel) Iterator(ctx context.Context, row types.Row) (types.RowIterator, error) { + return &dropModelRowIter{ + planner: p.planner, + ifExists: p.ifExists, + modelName: p.modelName, + }, nil +} + +func (p *PlanOpDropModel) WithChildren(children ...types.PlanOperator) (types.PlanOperator, error) { + return nil, nil +} + +type dropModelRowIter struct { + planner *ExecutionPlanner + ifExists bool + modelName string +} + +var _ types.RowIterator = (*dropModelRowIter)(nil) + +func (i *dropModelRowIter) Next(ctx context.Context) (types.Row, error) { + err := i.planner.checkAccess(ctx, i.modelName, accessTypeDropObject) + if err != nil { + return nil, err + } + + // check in the models table to see if it exists + v, err := i.planner.getModelByName(i.modelName) + if err != nil { + return nil, err + } + if v == nil { + if i.ifExists { + return nil, types.ErrNoMoreRows + } + return nil, sql3.NewErrModelNotFound(0, 0, i.modelName) + } + + err = i.planner.deleteModel(i.modelName) + if err != nil { + return nil, err + } + + return nil, types.ErrNoMoreRows +} diff --git a/sql3/planner/oppredict.go b/sql3/planner/oppredict.go new file mode 100644 index 000000000..ac9e4f233 --- /dev/null +++ b/sql3/planner/oppredict.go @@ -0,0 +1,253 @@ +// Copyright 2022 Molecula Corp. All rights reserved. + +package planner + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "github.com/featurebasedb/featurebase/v3/pql" + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" + "github.com/featurebasedb/featurebase/v3/sql3/planner/types" + "github.com/sajari/regression" +) + +// PlanOpPredict is an operator for a PREDICT +type PlanOpPredict struct { + ChildOp types.PlanOperator + planner *ExecutionPlanner + model *modelSystemObject + warnings []string +} + +func NewPlanOpPredict(planner *ExecutionPlanner, model *modelSystemObject, child types.PlanOperator) *PlanOpPredict { + return &PlanOpPredict{ + ChildOp: child, + planner: planner, + model: model, + warnings: make([]string, 0), + } +} + +func (p *PlanOpPredict) Schema() types.Schema { + result := make(types.Schema, 0) + + switch strings.ToLower(p.model.modelType) { + case "linear_regresssion": + labelName := p.model.labels[0] + + result = append(result, &types.PlannerColumn{ + ColumnName: fmt.Sprintf("predicted_%s", labelName), + RelationName: "", + AliasName: "", + // we need to get this type from somewhere...probably needs to be stored in the model def + Type: &parser.DataTypeDecimal{ + Scale: 4, + }, + }) + + default: + // don't add anything + } + // add the columns from the select + result = append(result, p.ChildOp.Schema()...) + + return result +} + +func (p *PlanOpPredict) Iterator(ctx context.Context, row types.Row) (types.RowIterator, error) { + // get the query iterator + iter, err := p.ChildOp.Iterator(ctx, row) + if err != nil { + return nil, err + } + + switch strings.ToLower(p.model.modelType) { + case "linear_regresssion": + return newLinearRegressionPredictIter(p.planner, p.model, p.ChildOp.Schema(), iter), nil + + default: + return nil, sql3.NewErrInternalf("unexpected model tyoe '%s'", p.model.modelType) + } +} + +func (p *PlanOpPredict) Children() []types.PlanOperator { + return []types.PlanOperator{ + p.ChildOp, + } +} + +func (p *PlanOpPredict) WithChildren(children ...types.PlanOperator) (types.PlanOperator, error) { + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return NewPlanOpPredict(p.planner, p.model, children[0]), nil +} + +func (p *PlanOpPredict) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_op"] = fmt.Sprintf("%T", p) + sc := make([]string, 0) + for _, e := range p.Schema() { + sc = append(sc, fmt.Sprintf("'%s', '%s', '%s'", e.ColumnName, e.RelationName, e.Type.TypeDescription())) + } + result["_schema"] = sc + result["model"] = p.model.name + result["child"] = p.ChildOp.Plan() + return result +} + +func (p *PlanOpPredict) String() string { + return "" +} + +func (p *PlanOpPredict) AddWarning(warning string) { + p.warnings = append(p.warnings, warning) +} + +func (p *PlanOpPredict) Warnings() []string { + var w []string + w = append(w, p.warnings...) + if p.ChildOp != nil { + w = append(w, p.ChildOp.Warnings()...) + } + return w +} + +type linearRegressionPredictIter struct { + child types.RowIterator + planner *ExecutionPlanner + model *modelSystemObject + regres *regression.Regression + childSchema types.Schema + inferenceRefs []*qualifiedRefPlanExpression + hasStarted *struct{} +} + +func newLinearRegressionPredictIter(planner *ExecutionPlanner, model *modelSystemObject, childSchema types.Schema, child types.RowIterator) *linearRegressionPredictIter { + return &linearRegressionPredictIter{ + planner: planner, + model: model, + child: child, + childSchema: childSchema, + regres: new(regression.Regression), + inferenceRefs: make([]*qualifiedRefPlanExpression, 0), + } +} + +func (i *linearRegressionPredictIter) Next(ctx context.Context) (types.Row, error) { + if i.hasStarted == nil { + // label column + i.regres.SetObserved("Murders per annum per 1,000,000 inhabitants") + // input columns + i.regres.SetVar(0, "Inhabitants") + i.regres.SetVar(1, "Percent with incomes below $5000") + i.regres.SetVar(2, "Percent unemployed") + + // go get the 'training set' from fb_model_data + iter := &tableScanRowIter{ + planner: i.planner, + tableName: "fb_model_data", + columns: []string{ + "_id", + "model_id", + "data", + }, + predicate: newBinOpPlanExpression( + newQualifiedRefPlanExpression("fb_model_data", "model_id", 0, parser.NewDataTypeString()), + parser.EQ, + newStringLiteralPlanExpression(i.model.name), + parser.NewDataTypeBool(), + ), + } + + for { + row, err := iter.Next(context.Background()) + if err != nil { + if err == types.ErrNoMoreRows { + break + } + return nil, err + } + fdata := make([]float64, 0) + err = json.Unmarshal([]byte(row[2].(string)), &fdata) + if err != nil { + return nil, err + } + label := fdata[0] + vars := fdata[1:] + i.regres.Train(regression.DataPoint(label, vars)) + } + + // run the regression + err := i.regres.Run() + if err != nil { + return nil, err + } + + fmt.Printf("Regression formula:\n%v\n", i.regres.Formula) + fmt.Printf("Regression:\n%s\n", i.regres) + + for _, ic := range i.model.inputColumns { + found := false + for j, s := range i.childSchema { + if strings.EqualFold(ic, s.ColumnName) { + if !typesAreAssignmentCompatible(parser.NewDataTypeDecimal(4), s.Type) { + return nil, sql3.NewErrInternalf("types not assignment compatible") + } + i.inferenceRefs = append(i.inferenceRefs, newQualifiedRefPlanExpression("", ic, j, s.Type)) + found = true + break + } + } + if !found { + return nil, sql3.NewErrInternalf("input column found found") + } + } + + i.hasStarted = &struct{}{} + } + + childrow, err := i.child.Next(ctx) + if err != nil { + return nil, err + } + + // construct the inference data + inferenceData := make([]float64, len(i.inferenceRefs)) + for j, ref := range i.inferenceRefs { + + val, err := ref.Evaluate(childrow) + if err != nil { + return nil, err + } + cval, err := coerceValue(ref.dataType, parser.NewDataTypeDecimal(4), val, parser.Pos{Line: 0, Column: 0}) + if err != nil { + return nil, err + } + dval := cval.(pql.Decimal) + inferenceData[j] = dval.Float64() + } + + // do the prediction + prediction, err := i.regres.Predict(inferenceData) + if err != nil { + return nil, err + } + + // turn the predition into a decimal + dprediction, err := pql.FromFloat64WithScale(prediction, 4) + if err != nil { + return nil, err + } + + // make an output row + row := make(types.Row, len(childrow)+1) + row[0] = dprediction + copy(row[1:], childrow) + + return row, nil +} diff --git a/sql3/planner/opsystemtable.go b/sql3/planner/opsystemtable.go index c6c2c74f4..ec21f5584 100644 --- a/sql3/planner/opsystemtable.go +++ b/sql3/planner/opsystemtable.go @@ -9,6 +9,7 @@ import ( "sort" pilosa "github.com/featurebasedb/featurebase/v3" + "github.com/featurebasedb/featurebase/v3/dax" "github.com/featurebasedb/featurebase/v3/pql" "github.com/featurebasedb/featurebase/v3/sql3" "github.com/featurebasedb/featurebase/v3/sql3/parser" @@ -500,6 +501,107 @@ type fbTableDDLRowIter struct { var _ types.RowIterator = (*fbTableDDLRowIter)(nil) +func generateTableDDL(tbl *dax.Table, newName string) string { + + var buf bytes.Buffer + buf.WriteString("create table ") + if len(newName) > 0 { + fmt.Fprintf(&buf, "%s", newName) + } else { + fmt.Fprintf(&buf, "%s", tbl.Name) + } + buf.WriteString(" (") + + for idx, col := range tbl.Fields { + if idx > 0 { + buf.WriteString(", ") + } + fmt.Fprintf(&buf, "%s", col.Name) + dataType := fieldSQLDataType(pilosa.FieldToFieldInfo(col)) + fmt.Fprintf(&buf, " %s", dataType.TypeDescription()) + + switch dt := dataType.(type) { + case *parser.DataTypeID, *parser.DataTypeString: + if col.Options.CacheType != pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { + fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) + } + if col.Options.CacheSize != pilosa.DefaultCacheSize && col.Options.CacheSize > 0 { + // if we still have the default, we need to print that out if we have a non-default size + if col.Options.CacheType == pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { + fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) + } + fmt.Fprintf(&buf, " size %d", col.Options.CacheSize) + } + + case *parser.DataTypeIDSet, *parser.DataTypeStringSet: + if col.Options.CacheType != pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { + fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) + } + if col.Options.CacheSize != pilosa.DefaultCacheSize && col.Options.CacheSize > 0 { + // if we still have the default, we need to print that out if we have a non-default size + if col.Options.CacheType == pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { + fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) + } + fmt.Fprintf(&buf, " size %d", col.Options.CacheSize) + } + + case *parser.DataTypeIDSetQuantum, *parser.DataTypeStringSetQuantum: + if col.Options.CacheType != pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { + fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) + } + if col.Options.CacheSize != pilosa.DefaultCacheSize && col.Options.CacheSize > 0 { + // if we still have the default, we need to print that out if we have a non-default size + if col.Options.CacheType == pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { + fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) + } + fmt.Fprintf(&buf, " size %d", col.Options.CacheSize) + } + if !col.Options.TimeQuantum.IsEmpty() { + fmt.Fprintf(&buf, " timequantum '%s'", col.Options.TimeQuantum) + } + if col.Options.TTL > 0 { + fmt.Fprintf(&buf, " ttl '%s'", col.Options.TTL.String()) + } + + case *parser.DataTypeInt: + minValue, maxValue := pql.MinMax(0) + + min := col.Options.Min + if !min.EqualTo(minValue) { + fmt.Fprintf(&buf, " min %d", min.ToInt64(0)) + } + + max := col.Options.Max + if !max.EqualTo(maxValue) { + fmt.Fprintf(&buf, " max %d", max.ToInt64(0)) + } + + case *parser.DataTypeDecimal: + minValue, maxValue := pql.MinMax(dt.Scale) + + min := col.Options.Min + if !min.EqualTo(minValue) { + fmt.Fprintf(&buf, " min %v", min) + } + + max := col.Options.Max + if !max.EqualTo(maxValue) { + fmt.Fprintf(&buf, " max %v", max) + } + + case *parser.DataTypeTimestamp: + if len(col.Options.TimeUnit) > 0 { + fmt.Fprintf(&buf, " timeunit '%s'", col.Options.TimeUnit) + } + // TODO(pok) how do we get epoch out of col? + + } + } + buf.WriteString(");") + + return buf.String() +} + func (i *fbTableDDLRowIter) Next(ctx context.Context) (types.Row, error) { if i.result == nil { tbls, err := i.planner.schemaAPI.Tables(ctx) @@ -512,98 +614,7 @@ func (i *fbTableDDLRowIter) Next(ctx context.Context) (types.Row, error) { for idx, tbl := range tbls { // build the ddl for this table - var buf bytes.Buffer - buf.WriteString("create table ") - fmt.Fprintf(&buf, "%s", tbl.Name) - buf.WriteString(" (") - - for idx, col := range tbl.Fields { - if idx > 0 { - buf.WriteString(", ") - } - fmt.Fprintf(&buf, "%s", col.Name) - dataType := fieldSQLDataType(pilosa.FieldToFieldInfo(col)) - fmt.Fprintf(&buf, " %s", dataType.TypeDescription()) - - switch dt := dataType.(type) { - case *parser.DataTypeID, *parser.DataTypeString: - if col.Options.CacheType != pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { - fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) - } - if col.Options.CacheSize != pilosa.DefaultCacheSize && col.Options.CacheSize > 0 { - // if we still have the default, we need to print that out if we have a non-default size - if col.Options.CacheType == pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { - fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) - } - fmt.Fprintf(&buf, " size %d", col.Options.CacheSize) - } - - case *parser.DataTypeIDSet, *parser.DataTypeStringSet: - if col.Options.CacheType != pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { - fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) - } - if col.Options.CacheSize != pilosa.DefaultCacheSize && col.Options.CacheSize > 0 { - // if we still have the default, we need to print that out if we have a non-default size - if col.Options.CacheType == pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { - fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) - } - fmt.Fprintf(&buf, " size %d", col.Options.CacheSize) - } - - case *parser.DataTypeIDSetQuantum, *parser.DataTypeStringSetQuantum: - if col.Options.CacheType != pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { - fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) - } - if col.Options.CacheSize != pilosa.DefaultCacheSize && col.Options.CacheSize > 0 { - // if we still have the default, we need to print that out if we have a non-default size - if col.Options.CacheType == pilosa.DefaultCacheType && len(col.Options.CacheType) > 0 { - fmt.Fprintf(&buf, " cachetype %s", col.Options.CacheType) - } - fmt.Fprintf(&buf, " size %d", col.Options.CacheSize) - } - if !col.Options.TimeQuantum.IsEmpty() { - fmt.Fprintf(&buf, " timequantum '%s'", col.Options.TimeQuantum) - } - if col.Options.TTL > 0 { - fmt.Fprintf(&buf, " ttl '%s'", col.Options.TTL.String()) - } - - case *parser.DataTypeInt: - minValue, maxValue := pql.MinMax(0) - - min := col.Options.Min - if !min.EqualTo(minValue) { - fmt.Fprintf(&buf, " min %d", min.ToInt64(0)) - } - - max := col.Options.Max - if !max.EqualTo(maxValue) { - fmt.Fprintf(&buf, " max %d", max.ToInt64(0)) - } - - case *parser.DataTypeDecimal: - minValue, maxValue := pql.MinMax(dt.Scale) - - min := col.Options.Min - if !min.EqualTo(minValue) { - fmt.Fprintf(&buf, " min %v", min) - } - - max := col.Options.Max - if !max.EqualTo(maxValue) { - fmt.Fprintf(&buf, " max %v", max) - } - - case *parser.DataTypeTimestamp: - if len(col.Options.TimeUnit) > 0 { - fmt.Fprintf(&buf, " timeunit '%s'", col.Options.TimeUnit) - } - // TODO(pok) how do we get epoch out of col? - - } - } - buf.WriteString(");") - ddl := buf.String() + ddl := generateTableDDL(tbl, "") i.result[idx] = &fbTableDDLRow{ id: string(tbl.Name), diff --git a/sql3/planner/planoptimizer.go b/sql3/planner/planoptimizer.go index 6733b1510..ffae15cc4 100644 --- a/sql3/planner/planoptimizer.go +++ b/sql3/planner/planoptimizer.go @@ -1,4 +1,4 @@ -// Copyright 2021 Molecula Corp. All rights reserved. +// Copyright 2023 Molecula Corp. All rights reserved. package planner @@ -14,6 +14,8 @@ import ( "github.com/featurebasedb/featurebase/v3/sql3/planner/types" ) +//TODO(pok) give every expression an id 'Expr1234' and use that to match on +//TODO(pok) have a rule to eliminate PlanOpRelAlias //TODO(pok) push filter down into join condition if terms reference either side of join //TODO(pok) push order by down as far as possible //TODO(pok) you can't group by _id in PQL, so we need to not use a PQL group by operator here @@ -707,6 +709,10 @@ func tryToReplaceGroupByWithPQLAggregate(ctx context.Context, a *ExecutionPlanne } thisNode.Aggregates[i] = newAgg + // these two can't be done in PQL + case *corrPlanExpression, *varPlanExpression: + return thisNode, true, nil + case types.Aggregable: switch ref := aggregable.FirstChildExpr().(type) { case *qualifiedRefPlanExpression: diff --git a/sql3/planner/systemobjects.go b/sql3/planner/systemobjects.go index 8fedda5bd..86a05393f 100644 --- a/sql3/planner/systemobjects.go +++ b/sql3/planner/systemobjects.go @@ -4,6 +4,7 @@ package planner import ( "context" + "encoding/json" "time" pilosa "github.com/featurebasedb/featurebase/v3" @@ -18,6 +19,20 @@ type viewSystemObject struct { statement string } +type functionSystemObject struct { + name string + language string + body string +} + +type modelSystemObject struct { + name string + status string + modelType string + labels []string + inputColumns []string +} + func (p *ExecutionPlanner) ensureViewsSystemTableExists(ctx context.Context) error { _, err := p.schemaAPI.TableByName(ctx, "fb_views") if err != nil { @@ -178,8 +193,8 @@ func (p *ExecutionPlanner) insertView(ctx context.Context, view *viewSystemObjec newStringLiteralPlanExpression(view.statement), newStringLiteralPlanExpression(""), newStringLiteralPlanExpression(""), - newDateLiteralPlanExpression(createTime), - newDateLiteralPlanExpression(createTime), + newTimestampLiteralPlanExpression(createTime), + newTimestampLiteralPlanExpression(createTime), }, }, } @@ -212,7 +227,7 @@ func (p *ExecutionPlanner) updateView(ctx context.Context, view *viewSystemObjec newStringLiteralPlanExpression(view.name), newStringLiteralPlanExpression(view.statement), newStringLiteralPlanExpression(""), - newDateLiteralPlanExpression(updateTime), + newTimestampLiteralPlanExpression(updateTime), }, }, } @@ -245,3 +260,621 @@ func (p *ExecutionPlanner) deleteView(ctx context.Context, viewName string) erro } return nil } + +func (p *ExecutionPlanner) ensureFunctionsSystemTableExists() error { + _, err := p.schemaAPI.TableByName(context.Background(), "fb_functions") + if err != nil { + if !isTableNotFoundError(err) { + return err + } + + // create table fb_functions ( + // _id string + // name string + // language string + // body string + // owner string + // updated_by string + // created_at timestamp + // updated_at timestamp + // ); + + // if it doesn't, create it by making the appropriate iterator + iter := &createTableRowIter{ + planner: p, + tableName: "fb_functions", + failIfExists: false, + isKeyed: true, + keyPartitions: 0, + columns: []*createTableField{ + { + planner: p, + name: "name", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "language", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "body", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "owner", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "updated_by", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "created_at", + typeName: dax.BaseTypeTimestamp, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeTimestamp(pilosa.DefaultEpoch, pilosa.TimeUnitSeconds), + }, + }, + { + planner: p, + name: "updated_at", + typeName: dax.BaseTypeTimestamp, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeTimestamp(pilosa.DefaultEpoch, pilosa.TimeUnitSeconds), + }, + }, + }, + description: "system table for functions", + } + // call next on our iterator to create the table + _, err := iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + } + return nil +} + +func (p *ExecutionPlanner) getFunctionByName(name string) (*functionSystemObject, error) { + err := p.ensureFunctionsSystemTableExists() + if err != nil { + return nil, err + } + + tbl, err := p.schemaAPI.TableByName(context.Background(), "fb_functions") + if err != nil { + return nil, sql3.NewErrTableNotFound(0, 0, "fb_functions") + } + + cols := make([]string, len(tbl.Fields)) + for i, c := range tbl.Fields { + cols[i] = string(c.Name) + } + + iter := &tableScanRowIter{ + planner: p, + tableName: "fb_functions", + columns: cols, + predicate: newBinOpPlanExpression( + newQualifiedRefPlanExpression("fb_functions", "_id", 0, parser.NewDataTypeString()), + parser.EQ, + newStringLiteralPlanExpression(name), + parser.NewDataTypeBool(), + ), + topExpr: nil, + } + + row, err := iter.Next(context.Background()) + if err != nil { + if err == types.ErrNoMoreRows { + // view does not exist + return nil, nil + } + return nil, err + } + + return &functionSystemObject{ + name: row[1].(string), + language: row[2].(string), + body: row[3].(string), + }, nil +} + +func (p *ExecutionPlanner) insertFunction(function *functionSystemObject) error { + err := p.ensureFunctionsSystemTableExists() + if err != nil { + return err + } + + createTime := time.Now().UTC() + + iter := &insertRowIter{ + planner: p, + tableName: "fb_functions", + targetColumns: []*qualifiedRefPlanExpression{ + newQualifiedRefPlanExpression("fb_functions", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "name", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "language", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "body", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "owner", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "updated_by", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "created_at", 0, parser.NewDataTypeTimestamp()), + newQualifiedRefPlanExpression("fb_functions", "updated_at", 0, parser.NewDataTypeTimestamp()), + }, + insertValues: [][]types.PlanExpression{ + { + newStringLiteralPlanExpression(function.name), + newStringLiteralPlanExpression(function.name), + newStringLiteralPlanExpression(function.language), + newStringLiteralPlanExpression(function.body), + newStringLiteralPlanExpression(""), + newStringLiteralPlanExpression(""), + newTimestampLiteralPlanExpression(createTime), + newTimestampLiteralPlanExpression(createTime), + }, + }, + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + return nil +} + +func (p *ExecutionPlanner) updateFunction(function *functionSystemObject) error { + err := p.ensureFunctionsSystemTableExists() + if err != nil { + return err + } + + updateTime := time.Now().UTC() + + iter := &insertRowIter{ + planner: p, + tableName: "fb_functions", + targetColumns: []*qualifiedRefPlanExpression{ + newQualifiedRefPlanExpression("fb_functions", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "language", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "body", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "updated_by", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_functions", "updated_at", 0, parser.NewDataTypeTimestamp()), + }, + insertValues: [][]types.PlanExpression{ + { + newStringLiteralPlanExpression(function.name), + newStringLiteralPlanExpression(function.language), + newStringLiteralPlanExpression(function.body), + newStringLiteralPlanExpression(""), + newTimestampLiteralPlanExpression(updateTime), + }, + }, + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + return nil +} + +func (p *ExecutionPlanner) deleteFunction(functionName string) error { + err := p.ensureFunctionsSystemTableExists() + if err != nil { + return err + } + + iter := &filteredDeleteRowIter{ + planner: p, + tableName: "fb_functions", + filter: newBinOpPlanExpression( + newQualifiedRefPlanExpression("fb_functions", "_id", 0, parser.NewDataTypeString()), + parser.EQ, + newStringLiteralPlanExpression(functionName), + parser.NewDataTypeBool(), + ), + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + return nil +} + +func (p *ExecutionPlanner) ensureModelsSystemTableExists() error { + _, err := p.schemaAPI.TableByName(context.Background(), "fb_models") + if err != nil { + if !isTableNotFoundError(err) { + return err + } + // create table fb_models ( + // _id string + // name string + // status string + // model_type string + // labels string --this is an array of string, we'll store it as a json object until we have string[] type in sql + // input_columns string --this is an array for string, we'll store it as a json object until we have string[] type in sql + // owner string + // updated_by string + // created_at timestamp + // updated_at timestamp + // ); + + // if it doesn't, create it by making the appropriate iterator + iter := &createTableRowIter{ + planner: p, + tableName: "fb_models", + failIfExists: false, + isKeyed: true, + keyPartitions: 0, + columns: []*createTableField{ + { + planner: p, + name: "name", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "status", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "model_type", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "labels", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "input_columns", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "owner", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "updated_by", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "created_at", + typeName: dax.BaseTypeTimestamp, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeTimestamp(pilosa.DefaultEpoch, pilosa.TimeUnitSeconds), + }, + }, + { + planner: p, + name: "updated_at", + typeName: dax.BaseTypeTimestamp, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeTimestamp(pilosa.DefaultEpoch, pilosa.TimeUnitSeconds), + }, + }, + }, + description: "system table for models", + } + // call next on our iterator to create the table + _, err := iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + } + return nil +} + +func (p *ExecutionPlanner) ensureModelDataSystemTableExists() error { + _, err := p.schemaAPI.TableByName(context.Background(), "fb_model_data") + if err != nil { + if !isTableNotFoundError(err) { + return err + } + // create table fb_model_data ( + // _id string + // model_id string + // data string + // ); + + // if it doesn't, create it by making the appropriate iterator + iter := &createTableRowIter{ + planner: p, + tableName: "fb_model_data", + failIfExists: false, + isKeyed: true, + keyPartitions: 0, + columns: []*createTableField{ + { + planner: p, + name: "model_id", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + { + planner: p, + name: "data", + typeName: dax.BaseTypeString, + fos: []pilosa.FieldOption{ + pilosa.OptFieldTypeMutex(pilosa.DefaultCacheType, pilosa.DefaultCacheSize), + pilosa.OptFieldKeys(), + }, + }, + }, + description: "system table for model data", + } + // call next on our iterator to create the table + _, err := iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + } + return nil +} + +func (p *ExecutionPlanner) getModelByName(name string) (*modelSystemObject, error) { + err := p.ensureModelsSystemTableExists() + if err != nil { + return nil, err + } + + tbl, err := p.schemaAPI.TableByName(context.Background(), "fb_models") + if err != nil { + return nil, sql3.NewErrTableNotFound(0, 0, "fb_models") + } + + cols := make([]string, len(tbl.Fields)) + for i, c := range tbl.Fields { + cols[i] = string(c.Name) + } + + iter := &tableScanRowIter{ + planner: p, + tableName: "fb_models", + columns: cols, + predicate: newBinOpPlanExpression( + newQualifiedRefPlanExpression("fb_models", "_id", 0, parser.NewDataTypeString()), + parser.EQ, + newStringLiteralPlanExpression(name), + parser.NewDataTypeBool(), + ), + topExpr: nil, + } + + row, err := iter.Next(context.Background()) + if err != nil { + if err == types.ErrNoMoreRows { + // model does not exist + return nil, nil + } + return nil, err + } + + labels := make([]string, 0) + err = json.Unmarshal([]byte(row[4].(string)), &labels) + if err != nil { + return nil, err + } + + inputColumns := make([]string, 0) + err = json.Unmarshal([]byte(row[5].(string)), &inputColumns) + if err != nil { + return nil, err + } + + return &modelSystemObject{ + name: row[1].(string), + status: row[2].(string), + modelType: row[3].(string), + labels: labels, + inputColumns: inputColumns, + }, nil +} + +func (p *ExecutionPlanner) insertModel(model *modelSystemObject) error { + err := p.ensureModelsSystemTableExists() + if err != nil { + return err + } + + createTime := time.Now().UTC() + + labelJson, err := json.Marshal(model.labels) + if err != nil { + return err + } + + inputColumnsJson, err := json.Marshal(model.inputColumns) + if err != nil { + return err + } + + iter := &insertRowIter{ + planner: p, + tableName: "fb_models", + targetColumns: []*qualifiedRefPlanExpression{ + newQualifiedRefPlanExpression("fb_models", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "name", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "status", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "model_type", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "labels", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "input_columns", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "owner", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "updated_by", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "created_at", 0, parser.NewDataTypeTimestamp()), + newQualifiedRefPlanExpression("fb_models", "updated_at", 0, parser.NewDataTypeTimestamp()), + }, + insertValues: [][]types.PlanExpression{ + { + newStringLiteralPlanExpression(model.name), + newStringLiteralPlanExpression(model.name), + newStringLiteralPlanExpression(model.status), + newStringLiteralPlanExpression(model.modelType), + newStringLiteralPlanExpression(string(labelJson)), + newStringLiteralPlanExpression(string(inputColumnsJson)), + newStringLiteralPlanExpression(""), + newStringLiteralPlanExpression(""), + newTimestampLiteralPlanExpression(createTime), + newTimestampLiteralPlanExpression(createTime), + }, + }, + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + return nil +} + +func (p *ExecutionPlanner) updateModel(model *modelSystemObject) error { + err := p.ensureModelsSystemTableExists() + if err != nil { + return err + } + + updateTime := time.Now().UTC() + + labelJson, err := json.Marshal(model.labels) + if err != nil { + return err + } + + inputColumnsJson, err := json.Marshal(model.inputColumns) + if err != nil { + return err + } + + iter := &insertRowIter{ + planner: p, + tableName: "fb_models", + targetColumns: []*qualifiedRefPlanExpression{ + newQualifiedRefPlanExpression("fb_models", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "status", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "model_type", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "labels", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "input_columns", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "updated_by", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_models", "updated_at", 0, parser.NewDataTypeTimestamp()), + }, + insertValues: [][]types.PlanExpression{ + { + newStringLiteralPlanExpression(model.name), + newStringLiteralPlanExpression(model.status), + newStringLiteralPlanExpression(model.modelType), + newStringLiteralPlanExpression(string(labelJson)), + newStringLiteralPlanExpression(string(inputColumnsJson)), + newStringLiteralPlanExpression(""), + newTimestampLiteralPlanExpression(updateTime), + }, + }, + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + return nil +} + +func (p *ExecutionPlanner) deleteModel(modelName string) error { + err := p.ensureModelsSystemTableExists() + if err != nil { + return err + } + + err = p.ensureModelDataSystemTableExists() + if err != nil { + return err + } + + iter := &filteredDeleteRowIter{ + planner: p, + tableName: "fb_models", + filter: newBinOpPlanExpression( + newQualifiedRefPlanExpression("fb_models", "_id", 0, parser.NewDataTypeString()), + parser.EQ, + newStringLiteralPlanExpression(modelName), + parser.NewDataTypeBool(), + ), + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + + iter = &filteredDeleteRowIter{ + planner: p, + tableName: "fb_model_data", + filter: newBinOpPlanExpression( + newQualifiedRefPlanExpression("fb_model_data", "model_id", 0, parser.NewDataTypeString()), + parser.EQ, + newStringLiteralPlanExpression(modelName), + parser.NewDataTypeBool(), + ), + } + _, err = iter.Next(context.Background()) + if err != nil && err != types.ErrNoMoreRows { + return err + } + + return nil +} diff --git a/sql3/planner/userdefinedfunctions.go b/sql3/planner/userdefinedfunctions.go new file mode 100644 index 000000000..8b0935c08 --- /dev/null +++ b/sql3/planner/userdefinedfunctions.go @@ -0,0 +1,71 @@ +package planner + +import ( + "github.com/featurebasedb/featurebase/v3/sql3" + "github.com/featurebasedb/featurebase/v3/sql3/parser" +) + +func (p *ExecutionPlanner) analyzeUserDefinedFunction(call *parser.Call, scope parser.Statement, function *functionSystemObject) (parser.Expr, error) { + // TODO(pok) removing user defined functions for now + // // hard code to 1 string parameter + // if len(call.Args) != 1 { + // return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 2, len(call.Args)) + // } + + // // arg 1 + // argType := parser.NewDataTypeString() + // if !typesAreAssignmentCompatible(argType, call.Args[0].DataType()) { + // return nil, sql3.NewErrParameterTypeMistmatch(call.Args[0].Pos().Line, call.Args[0].Pos().Column, call.Args[0].DataType().TypeDescription(), argType.TypeDescription()) + // } + + // //return string + // call.ResultDataType = parser.NewDataTypeString() + + // return call, nil + return nil, sql3.NewErrUnsupported(0, 0, false, "user defined functions") +} + +func (n *callPlanExpression) evaluateUserDefinedFunction(currentRow []interface{}) (interface{}, error) { + // TODO(pok) removing what effectively is a remote code exploit + // we will come back to this to add sql udfs and external code later + + // argEval, err := n.args[0].Evaluate(currentRow) + // if err != nil { + // return nil, err + // } + // // nil if anything is nil + // if argEval == nil { + // return nil, nil + // } + + // //get the value + // coercedArg, err := coerceValue(n.args[0].Type(), parser.NewDataTypeString(), argEval, parser.Pos{Line: 0, Column: 0}) + // if err != nil { + // return nil, err + // } + + // arg, argOk := coercedArg.(string) + // if !argOk { + // return nil, sql3.NewErrInternalf("unable to convert value") + // } + + // // save the body to a temp file + // file, err := os.CreateTemp("", "py-body") + // if err != nil { + // return nil, err + // } + // defer os.Remove(file.Name()) + + // file.Write([]byte(n.udfReference.body)) + + // cmd := exec.Command("python3", file.Name(), arg) + // stdout, err := cmd.Output() + + // if err != nil { + // return nil, err + // } + + // retVal := string(stdout) + // return retVal, nil + return nil, sql3.NewErrUnsupported(0, 0, false, "user defined functions") +} diff --git a/sql3/sql_complex_test.go b/sql3/sql_complex_test.go index 5d35d2c5c..05b55a65e 100644 --- a/sql3/sql_complex_test.go +++ b/sql3/sql_complex_test.go @@ -51,7 +51,7 @@ func TestPlanner_SystemTableFanout(t *testing.T) { server := c.GetNode(0).Server t.Run("PerfCounters", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, server, `select * from fb_performance_counters`) + results, columns, _, err := sql_test.MustQueryRows(t, nil, server, `select * from fb_performance_counters`) if err != nil { t.Fatal(err) } @@ -72,7 +72,7 @@ func TestPlanner_SystemTableFanout(t *testing.T) { }) t.Run("SystemTablesExecRequests", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select * from fb_exec_requests`) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select * from fb_exec_requests`) if err != nil { t.Fatal(err) } @@ -105,7 +105,7 @@ func TestPlanner_SystemTableFanout(t *testing.T) { }) t.Run("SystemTablesExecRequestsAgg", func(t *testing.T) { - _, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select + _, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select count(request_id) as request_count, min(elapsed_time) as min_duration, max(elapsed_time) as max_duration, @@ -127,6 +127,23 @@ func TestPlanner_SystemTableFanout(t *testing.T) { t.Fatal(diff) } }) + + // additional testing for *ExecutionPlanner.mapper is going here because + // this is where the existing coverage comes from. + t.Run("MapperContextCancellation", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + // give the query time to start, but not enough time to finish. + // works in conjunction with the 100 ms delay in a MustRunQuery + // that's been given a non-nil context. + go func() { + time.Sleep(50 * time.Millisecond) + cancel() + }() + _, _, _, err := sql_test.MustQueryRows(t, ctx, server, `select * from fb_performance_counters`) + if err != context.Canceled { + t.Fatalf("expected %v error, got %v", context.Canceled, err) + } + }) } func TestPlanner_Show(t *testing.T) { @@ -156,7 +173,7 @@ func TestPlanner_Show(t *testing.T) { } t.Run("SystemTablesInfo", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select name, platform, platform_version, db_version, state, node_count, replica_count from fb_database_info`) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select name, platform, platform_version, db_version, state, node_count, replica_count from fb_database_info`) if err != nil { t.Fatal(err) } @@ -178,7 +195,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("SystemTablesNode", func(t *testing.T) { - _, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select * from fb_database_nodes`) + _, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select * from fb_database_nodes`) if err != nil { t.Fatal(err) } @@ -197,7 +214,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("ShowDatabases", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `SHOW DATABASES`) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `SHOW DATABASES`) if err != nil { t.Fatal(err) } @@ -223,7 +240,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("ShowTables", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `SHOW TABLES`) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `SHOW TABLES`) if err != nil { t.Fatal(err) } @@ -250,7 +267,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("ShowCreateTable", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SHOW CREATE TABLE %i`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SHOW CREATE TABLE %i`, c)) if err != nil { t.Fatal(err) } @@ -272,7 +289,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("ShowCreateTableCacheTypes", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table iris1 ( + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table iris1 ( _id id, speciesid id cachetype ranked size 1000 species string cachetype ranked size 1000 @@ -287,7 +304,7 @@ func TestPlanner_Show(t *testing.T) { t.Fatal(err) } - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `SHOW CREATE TABLE iris1`) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `SHOW CREATE TABLE iris1`) if err != nil { t.Fatal(err) } @@ -309,7 +326,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("ShowColumns", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SHOW COLUMNS FROM %i`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SHOW COLUMNS FROM %i`, c)) if err != nil { t.Fatal(err) } @@ -338,7 +355,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("ShowColumns2", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SHOW COLUMNS FROM %l`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SHOW COLUMNS FROM %l`, c)) if err != nil { t.Fatal(err) } @@ -367,7 +384,7 @@ func TestPlanner_Show(t *testing.T) { }) t.Run("ShowColumnsFromNotATable", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `SHOW COLUMNS FROM foo`) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `SHOW COLUMNS FROM foo`) if err != nil { if err.Error() != "[1:19] table 'foo' not found" { t.Fatal(err) @@ -426,7 +443,7 @@ func TestPlanner_CoverCreateTable(t *testing.T) { sql += `) keypartitions 12` // Run the create table statement. - _, _, _, err := sql_test.MustQueryRows(t, server, sql) + _, _, _, err := sql_test.MustQueryRows(t, nil, server, sql) if assert.Error(t, err) { assert.Equal(t, fld.expErr, err.Error()) // sql3.SQLErrConflictingColumnConstraint.Message @@ -603,7 +620,7 @@ func TestPlanner_CoverCreateTable(t *testing.T) { sql += `) keypartitions 12` // Run the create table statement. - results, columns, _, err := sql_test.MustQueryRows(t, server, sql) + results, columns, _, err := sql_test.MustQueryRows(t, nil, server, sql) assert.NoError(t, err) assert.Equal(t, [][]interface{}{}, results) assert.Equal(t, []*pilosa.WireQueryField{}, columns) @@ -666,7 +683,7 @@ func TestPlanner_CreateTable(t *testing.T) { server := c.GetNode(0).Server t.Run("CreateTableAllDataTypes", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, server, `create table allcoltypes ( + results, columns, _, err := sql_test.MustQueryRows(t, nil, server, `create table allcoltypes ( _id id, intcol int, boolcol bool, @@ -689,7 +706,7 @@ func TestPlanner_CreateTable(t *testing.T) { }) t.Run("CreateTableAllDataTypesAgain", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, `create table allcoltypes ( + _, _, _, err := sql_test.MustQueryRows(t, nil, server, `create table allcoltypes ( _id id, intcol int, boolcol bool, @@ -709,28 +726,28 @@ func TestPlanner_CreateTable(t *testing.T) { }) t.Run("CreateTableMixedCaseColumn", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, `create table lowercase (_id id, name string, SomeColumn string, legalname string);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, server, `create table lowercase (_id id, name string, SomeColumn string, legalname string);`) if err != nil { t.Fatal(err) } }) t.Run("CreateTableMixedCaseColumn", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, `create table MixedCcase (_id id, name string, SomeColumn string, legalname string);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, server, `create table MixedCcase (_id id, name string, SomeColumn string, legalname string);`) if err != nil { t.Fatal(err) } }) t.Run("DropTable1", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, `drop table allcoltypes`) + _, _, _, err := sql_test.MustQueryRows(t, nil, server, `drop table allcoltypes`) if err != nil { t.Fatal(err) } }) t.Run("CreateTableAllDataTypesAllConstraints", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, server, `create table allcoltypes ( + results, columns, _, err := sql_test.MustQueryRows(t, nil, server, `create table allcoltypes ( _id id, intcol int min 0 max 10000, boolcol bool, @@ -756,7 +773,7 @@ func TestPlanner_CreateTable(t *testing.T) { }) t.Run("ShowColumns1", func(t *testing.T) { - _, columns, _, err := sql_test.MustQueryRows(t, server, `SHOW COLUMNS FROM allcoltypes`) + _, columns, _, err := sql_test.MustQueryRows(t, nil, server, `SHOW COLUMNS FROM allcoltypes`) if err != nil { t.Fatal(err) } @@ -781,7 +798,7 @@ func TestPlanner_CreateTable(t *testing.T) { }) t.Run("CreateTableDupeColumns", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, `create table dupecols ( + _, _, _, err := sql_test.MustQueryRows(t, nil, server, `create table dupecols ( _id id, _id int)`) if err == nil { @@ -794,7 +811,7 @@ func TestPlanner_CreateTable(t *testing.T) { }) t.Run("CreateTableMissingId", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, `create table missingid ( + _, _, _, err := sql_test.MustQueryRows(t, nil, server, `create table missingid ( foo int)`) if err == nil { t.Fatal("expected error") @@ -824,7 +841,7 @@ func TestPlanner_AlterTable(t *testing.T) { server := c.GetNode(0).Server t.Run("AlterTableDrop", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, server, fmt.Sprintf(`alter table %i drop column f`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, server, fmt.Sprintf(`alter table %i drop column f`, c)) if err != nil { t.Fatal(err) } @@ -838,7 +855,7 @@ func TestPlanner_AlterTable(t *testing.T) { }) t.Run("AlterTableAdd", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, server, fmt.Sprintf(`alter table %i add column f int`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, server, fmt.Sprintf(`alter table %i add column f int`, c)) if err != nil { t.Fatal(err) } @@ -852,7 +869,7 @@ func TestPlanner_AlterTable(t *testing.T) { }) // test for adding duplicate column, this test requires column 'f' to be defined beforehand. t.Run("AlterTableAdd", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, fmt.Sprintf(`alter table %i add column f int`, c)) + _, _, _, err := sql_test.MustQueryRows(t, nil, server, fmt.Sprintf(`alter table %i add column f int`, c)) if err == nil { t.Fatal("expected error") } else { @@ -863,7 +880,7 @@ func TestPlanner_AlterTable(t *testing.T) { }) // test for bad column definitions. t.Run("AlterTableAdd", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, fmt.Sprintf(`alter table %i add column dt date`, c)) + _, _, _, err := sql_test.MustQueryRows(t, nil, server, fmt.Sprintf(`alter table %i add column dt date`, c)) if err == nil { t.Fatal("expected error") } else { @@ -874,7 +891,7 @@ func TestPlanner_AlterTable(t *testing.T) { }) //test for the system rule that enforces the special primary key column "_id" can't be added using alter table statement. t.Run("AlterTableAdd", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, server, fmt.Sprintf(`alter table %i add column _id int`, c)) + _, _, _, err := sql_test.MustQueryRows(t, nil, server, fmt.Sprintf(`alter table %i add column _id int`, c)) if err == nil { t.Fatal("expected error") } else { @@ -885,7 +902,7 @@ func TestPlanner_AlterTable(t *testing.T) { }) t.Run("AlterTableRename", func(t *testing.T) { t.Skip("not yet implemented") - results, columns, _, err := sql_test.MustQueryRows(t, server, fmt.Sprintf(`alter table %i rename column f to g`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, server, fmt.Sprintf(`alter table %i rename column f to g`, c)) if err != nil { t.Fatal(err) } @@ -915,29 +932,29 @@ func TestPlanner_DropThings(t *testing.T) { } t.Run("DropTable", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`DROP TABLE %i`, c)) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`DROP TABLE %i`, c)) if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`DROP TABLE %j`, c)) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`DROP TABLE %j`, c)) if err == nil || !strings.Contains(err.Error(), `not found`) { t.Fatalf("expected 'table not found', got %v", err) } }) t.Run("DropView", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `CREATE VIEW vw AS SELECT true`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `CREATE VIEW vw AS SELECT true`) if err != nil { t.Fatalf("creating view: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `DROP VIEW vw`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `DROP VIEW vw`) if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `DROP VIEW vw`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `DROP VIEW vw`) if err == nil || !strings.Contains(err.Error(), `not found`) { t.Fatalf("expected 'table not found', got %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `DROP VIEW IF EXISTS vw`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `DROP VIEW IF EXISTS vw`) if err != nil { t.Fatal(err) } @@ -984,7 +1001,7 @@ func TestPlanner_ExpressionsInSelectListParen(t *testing.T) { } t.Run("ParenOne", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT (a != b) = false, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT (a != b) = false, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1005,7 +1022,7 @@ func TestPlanner_ExpressionsInSelectListParen(t *testing.T) { }) t.Run("ParenTwo", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT (a != b) = (false), _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT (a != b) = (false), _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1064,7 +1081,7 @@ func TestPlanner_ExpressionsInSelectListLiterals(t *testing.T) { } t.Run("LiteralsBool", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT false = true, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT false = true, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1085,7 +1102,7 @@ func TestPlanner_ExpressionsInSelectListLiterals(t *testing.T) { }) t.Run("LiteralsInt", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT 1 + 2, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT 1 + 2, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1106,7 +1123,7 @@ func TestPlanner_ExpressionsInSelectListLiterals(t *testing.T) { }) t.Run("LiteralsID", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT _id + 2, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT _id + 2, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1127,7 +1144,7 @@ func TestPlanner_ExpressionsInSelectListLiterals(t *testing.T) { }) t.Run("LiteralsDecimal", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT d + 2.0, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT d + 2.0, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1152,7 +1169,7 @@ func TestPlanner_ExpressionsInSelectListLiterals(t *testing.T) { }) t.Run("LiteralsString", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT str || ' bar', _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT str || ' bar', _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1211,7 +1228,7 @@ func TestPlanner_ExpressionsInSelectListCase(t *testing.T) { } t.Run("CaseWithBase", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT b, case b when 100 then 10 when 201 then 20 else 5 end, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT b, case b when 100 then 10 when 201 then 20 else 5 end, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1233,7 +1250,7 @@ func TestPlanner_ExpressionsInSelectListCase(t *testing.T) { }) t.Run("CaseWithNoBase", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT b, case when b = 100 then 10 when b = 201 then 20 else 5 end, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT b, case when b = 100 then 10 when b = 201 then 20 else 5 end, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1295,7 +1312,7 @@ func TestPlanner_Select(t *testing.T) { } t.Run("UnqualifiedColumns", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b, _id FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b, _id FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1317,7 +1334,7 @@ func TestPlanner_Select(t *testing.T) { }) t.Run("QualifiedTableRef", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT bar.a, bar.b, bar._id FROM %j as bar`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT bar.a, bar.b, bar._id FROM %j as bar`, c)) if err != nil { t.Fatal(err) } @@ -1339,7 +1356,7 @@ func TestPlanner_Select(t *testing.T) { }) t.Run("AliasedUnqualifiedColumns", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT a as foo, b as bar, _id as baz FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT a as foo, b as bar, _id as baz FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1361,7 +1378,7 @@ func TestPlanner_Select(t *testing.T) { }) t.Run("QualifiedColumns", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT %j._id, %j.a, %j.b FROM %j`, c, c, c, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT %j._id, %j.a, %j.b FROM %j`, c, c, c, c)) if err != nil { t.Fatal(err) } @@ -1383,7 +1400,7 @@ func TestPlanner_Select(t *testing.T) { }) t.Run("UnqualifiedStar", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT * FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT * FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1405,7 +1422,7 @@ func TestPlanner_Select(t *testing.T) { }) t.Run("QualifiedStar", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT %j.* FROM %j`, c, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT %j.* FROM %j`, c, c)) if err != nil { t.Fatal(err) } @@ -1427,7 +1444,7 @@ func TestPlanner_Select(t *testing.T) { }) t.Run("NoIdentifier", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b FROM %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b FROM %j`, c)) if err != nil { t.Fatal(err) } @@ -1484,7 +1501,7 @@ func TestPlanner_SelectOrderBy(t *testing.T) { } t.Run("OrderBy", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b, _id FROM %j order by a desc`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b, _id FROM %j order by a desc`, c)) if err != nil { t.Fatal(err) } @@ -1510,22 +1527,22 @@ func TestPlanner_BulkInsert(t *testing.T) { c := test.MustRunCluster(t, 1) defer c.Close() - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, "create table j (_id id, a int, b int)") + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, "create table j (_id id, a int, b int)") if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, "create table j1 (_id id, a int, b int)") + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, "create table j1 (_id id, a int, b int)") if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, "create table j2 (_id id, a int, b int)") + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, "create table j2 (_id id, a int, b int)") if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table alltypes ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table alltypes ( _id id, id1 id, i1 int, @@ -1541,96 +1558,96 @@ func TestPlanner_BulkInsert(t *testing.T) { } t.Run("BulkBadMap", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0, 1 int, 2 int) from '/Users/bar/foo.csv';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0, 1 int, 2 int) from '/Users/bar/foo.csv';`) if err == nil || !strings.Contains(err.Error(), `expected type name, found ','`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkNoWith", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv';`) if err == nil || !strings.Contains(err.Error(), ` expected WITH, found ';'`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadWith", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH UNICORNS AND RAINBOWS;`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH UNICORNS AND RAINBOWS;`) if err == nil || !strings.Contains(err.Error(), `expected BATCHSIZE, ROWSLIMIT, FORMAT, INPUT, ALLOW_MISSING_VALUES or HEADER_ROW, found UNICORNS`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkNoWithFormat", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' with batchsize 2;`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' with batchsize 2;`) if err == nil || !strings.Contains(err.Error(), `format specifier expected`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadWithFormat", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'BLAH';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'BLAH';`) if err == nil || !strings.Contains(err.Error(), `invalid format specifier 'BLAH'`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkNoWithInput", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV';`) if err == nil || !strings.Contains(err.Error(), `input specifier expected`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadWithInput", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'WOOPWOOP';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'WOOPWOOP';`) if err == nil || !strings.Contains(err.Error(), `invalid input specifier 'WOOPWOOP'`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadTable", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into foo (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into foo (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `table 'foo' not found`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkNoID", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (a, b) map (0 int, 1 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (a, b) map (0 int, 1 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `insert column list must have '_id' column specified`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkNoNonID", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id) map (0 id) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id) map (0 id) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `insert column list must have at least one non '_id' column specified`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadColumn", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, k, l) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, k, l) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `column 'k' not found`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkMapCountMismatch", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `mismatch in the count of expressions and target columns`) { t.Fatalf("unexpected error: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int, 3 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int, 3 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `mismatch in the count of expressions and target columns`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkCSVFileNonExistent", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/Users/bar/foo.csv' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `unable to read datasource '/Users/bar/foo.csv': file '/Users/bar/foo.csv' does not exist`) { t.Fatalf("unexpected error: %v", err) } @@ -1652,19 +1669,19 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j1 (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE';`, tmpfile.Name())) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j1 (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE';`, tmpfile.Name())) if err == nil || !strings.Contains(err.Error(), `value '_id' cannot be converted to type 'id'`) { t.Fatalf("unexpected error: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j1 (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE' HEADER_ROW;`, tmpfile.Name())) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j1 (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE' HEADER_ROW;`, tmpfile.Name())) if err != nil { t.Fatal(err) } }) t.Run("BulkCSVBadMap", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 10 int) from x'1,10,20 + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 10 int) from x'1,10,20 2,11,21 3,12,22 4,13,23 @@ -1680,36 +1697,36 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkBadSource", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from 23 WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from 23 WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `string literal expected`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadFormat", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from 'foo' WITH FORMAT 12 INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from 'foo' WITH FORMAT 12 INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `string literal expected`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadMap", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, "3" int, 2 int) from 'foo' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, "3" int, 2 int) from 'foo' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `integer literal expected`) { t.Fatalf("unexpected error: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, "3" int, 2 int) from 'foo' WITH FORMAT 'NDJSON' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, "3" int, 2 int) from 'foo' WITH FORMAT 'NDJSON' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `string literal expected`) { t.Fatalf("unexpected error: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 3 aunt, 2 int) from 'foo' WITH FORMAT 'CSV' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 3 aunt, 2 int) from 'foo' WITH FORMAT 'CSV' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `unknown type 'aunt'`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadInput", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from 'foo' WITH FORMAT 'CSV' INPUT 23;`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from 'foo' WITH FORMAT 'CSV' INPUT 23;`) if err == nil || !strings.Contains(err.Error(), `string literal expected`) { t.Fatalf("unexpected error: %v", err) } @@ -1731,7 +1748,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE';`, tmpfile.Name())) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE';`, tmpfile.Name())) if err != nil { t.Fatal(err) } @@ -1753,42 +1770,42 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE';`, tmpfile.Name())) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE';`, tmpfile.Name())) if err != nil { t.Fatal(err) } }) t.Run("BulkCSVFileBadBatchSize", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/foo/bar' WITH FORMAT 'CSV' INPUT 'FILE' BATCHSIZE 0;`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/foo/bar' WITH FORMAT 'CSV' INPUT 'FILE' BATCHSIZE 0;`) if err == nil || !strings.Contains(err.Error(), `invalid batch size '0'`) { t.Fatalf("unexpected error: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/foo/bar' WITH FORMAT 'CSV' INPUT 'FILE' BATCHSIZE 'foo';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/foo/bar' WITH FORMAT 'CSV' INPUT 'FILE' BATCHSIZE 'foo';`) if err == nil || !strings.Contains(err.Error(), `integer literal expected`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkNDJSONFileBadBatchSize", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('id' id, 'a' int, 'b' int) from '/foo/bar' WITH FORMAT 'NDJSON' INPUT 'FILE' BATCHSIZE 0;`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('id' id, 'a' int, 'b' int) from '/foo/bar' WITH FORMAT 'NDJSON' INPUT 'FILE' BATCHSIZE 0;`) if err == nil || !strings.Contains(err.Error(), `invalid batch size '0'`) { t.Fatalf("unexpected error: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('id' id, 'a' int, 'b' int) from '/foo/bar' WITH FORMAT 'NDJSON' INPUT 'FILE' BATCHSIZE 'foo';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('id' id, 'a' int, 'b' int) from '/foo/bar' WITH FORMAT 'NDJSON' INPUT 'FILE' BATCHSIZE 'foo';`) if err == nil || !strings.Contains(err.Error(), `integer literal expected`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkBadRowsLimit", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/foo/bar' WITH FORMAT 'CSV' INPUT 'FILE' ROWSLIMIT 'foo';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from '/foo/bar' WITH FORMAT 'CSV' INPUT 'FILE' ROWSLIMIT 'foo';`) if err == nil || !strings.Contains(err.Error(), `integer literal expected`) { t.Fatalf("unexpected error: %v", err) } }) t.Run("BulkTransformBadName", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) transform (@0, @1, @z) from 'foo' WITH FORMAT 'NDJSON' INPUT 'FILE';`) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) transform (@0, @1, @z) from 'foo' WITH FORMAT 'NDJSON' INPUT 'FILE';`) if err == nil || !strings.Contains(err.Error(), `unknown identifier 'z'`) { t.Fatalf("unexpected error: %v", err) } @@ -1810,12 +1827,12 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j2 (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE' ROWSLIMIT 2;`, tmpfile.Name())) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j2 (_id, a, b) map (0 id, 1 int, 2 int) from '%s' WITH FORMAT 'CSV' INPUT 'FILE' ROWSLIMIT 2;`, tmpfile.Name())) if err != nil { t.Fatal(err) } - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `SELECT count(*) from j2`) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `SELECT count(*) from j2`) if err != nil { t.Fatal(err) } @@ -1834,14 +1851,14 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkCSVBlobDefault", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, "bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from x'1,10,20\n2,11,21\n3,12,22\n4,13,23\n5,13,23\n6,13,23\n7,13,23\n8,13,23\n9,13,23\n10,13,23' WITH FORMAT 'CSV' INPUT 'STREAM';") + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, "bulk insert into j (_id, a, b) map (0 id, 1 int, 2 int) from x'1,10,20\n2,11,21\n3,12,22\n4,13,23\n5,13,23\n6,13,23\n7,13,23\n8,13,23\n9,13,23\n10,13,23' WITH FORMAT 'CSV' INPUT 'STREAM';") if err != nil { t.Fatal(err) } }) t.Run("BulkNDJsonBlobDefault", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) from x'{ "_id": 1, "a": 10, "b": 20 } { "_id": 2, "a": 10, "b": 20 } { "_id": 3, "a": 10, "b": 20 } @@ -1858,7 +1875,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkNDJsonBlobBadPath", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.frobny' int) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.frobny' int) from x'{ "_id": 1, "a": 10, "b": 20 } { "_id": 2, "a": 10, "b": 20 } { "_id": 3, "a": 10, "b": 20 } @@ -1891,7 +1908,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) from '%s' WITH FORMAT 'NDJSON' INPUT 'FILE';`, tmpfile.Name())) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) from '%s' WITH FORMAT 'NDJSON' INPUT 'FILE';`, tmpfile.Name())) if err != nil { t.Fatal(err) } @@ -1914,14 +1931,14 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) transform (@0, @1, @2) from '%s' WITH FORMAT 'NDJSON' INPUT 'FILE';`, tmpfile.Name())) + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) transform (@0, @1, @2) from '%s' WITH FORMAT 'NDJSON' INPUT 'FILE';`, tmpfile.Name())) if err != nil { t.Fatal(err) } }) t.Run("BulkNDJsonAllTypes", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into alltypes (_id, id1, i1, ids1, ss1, ts1, s1, b1, d1) map ('$._id' id, '$.id1' id, '$.i1' int, '$.ids1' idset, '$.ss1' stringset, '$.ts1' timestamp, '$.s1' string, '$.b1' bool, '$.d1' decimal(2)) @@ -1942,7 +1959,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkNDJsonBadJsonPath", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into alltypes (_id, id1, i1, ids1, ss1, ts1, s1, b1, d1) map ('$._id' id, '$.id1' id, '$.i1' int, '$.ids1' idset, '$.ss1' stringset, '$.ts1' timestamp, '$.s1' string, '$.blah' bool, '$.d1' decimal(2)) @@ -1961,7 +1978,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkNDJsonBadJson", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into alltypes (_id, id1, i1, ids1, ss1, ts1, s1, b1, d1) map ('$._id' id, '$.id1' id, '$.i1' int, '$.ids1' idset, '$.ss1' stringset, '$.ts1' timestamp, '$.s1' string, '$.b1' bool, '$.d1' decimal(2)) @@ -1980,7 +1997,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkInsertDecimals", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table iris ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table iris ( _id id, sepallength decimal(2), sepalwidth decimal(2), @@ -1992,7 +2009,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into iris (_id, sepallength, sepalwidth, petallength, petalwidth, species) map('id' id, 'sepalLength' DECIMAL, @@ -2011,7 +2028,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatalf("unexpected error: %v", err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into iris (_id, sepallength, sepalwidth, petallength, petalwidth, species) map('id' id, 'sepalLength' DECIMAL(2), @@ -2032,7 +2049,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkInsertDupeColumnPlusNullsInJson", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table dataviz (_id string, guid string, aba string,amount int, + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table dataviz (_id string, guid string, aba string,amount int, audit_id id, bools bool, bools-exist bool, browser string, browser_version string, central_group string, device string, error_description string, event_date_str string, event_epoch int, event_length int, event_type string, fidb string, gt_status string, gt_type string, operating_system string, os_version string, transaction_id string, user_id id);`) @@ -2040,7 +2057,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into dataviz (_id, aba, amount, audit_id, bools, bools-exist, browser, browser_version, central_group, device, error_description, event_date_str, event_epoch,event_length,event_type,fidb,gt_status,gt_type,_id,operating_system,os_version,transaction_id,user_id) map('guid' string,'aba' string, 'amount' int, 'audit_id' id, 'event_success' bool, 'db_success' bool, 'browser' string, 'browser_version' string, 'central_group' string, 'device' string, 'error_description' string, 'event_date_str' string, 'event_epoch' int, 'event_length' int, 'event_type' string, 'fidb' string, 'gt_status' string, 'gt_type' string, 'guid' string, 'operating_system' string, 'os_version' string, 'transaction_id' string, 'user_id' id) from @@ -2056,7 +2073,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkInsertCSVStringIDSet", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table greg-test ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table greg-test ( _id STRING, id_col ID, string_col STRING cachetype ranked size 1000, @@ -2071,7 +2088,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `BULK INSERT INTO greg-test ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `BULK INSERT INTO greg-test ( _id, id_col, string_col, @@ -2114,7 +2131,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkInsertCSVStringNulls", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table greg-test-n ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table greg-test-n ( _id ID, id_col ID, string_col STRING cachetype ranked size 1000, @@ -2129,7 +2146,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `BULK INSERT INTO greg-test-n ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `BULK INSERT INTO greg-test-n ( _id, id_col, string_col, @@ -2172,7 +2189,7 @@ func TestPlanner_BulkInsert(t *testing.T) { }) t.Run("BulkInsertAllowMissingValues", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table greg-test-amv ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table greg-test-amv ( _id STRING, id_col ID, string_col STRING cachetype ranked size 1000, @@ -2187,7 +2204,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `BULK INSERT INTO greg-test-amv ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `BULK INSERT INTO greg-test-amv ( _id, id_col, string_col, @@ -2229,7 +2246,7 @@ func TestPlanner_BulkInsert(t *testing.T) { } }) t.Run("BulkInsertNDJSONStringIDSet", func(t *testing.T) { - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table greg-test-01 ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table greg-test-01 ( _id STRING, id_col ID, string_col STRING cachetype ranked size 1000, @@ -2244,7 +2261,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `BULK INSERT INTO greg-test-01 ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `BULK INSERT INTO greg-test-01 ( _id, id_col, string_col, @@ -2282,7 +2299,7 @@ func TestPlanner_BulkInsert(t *testing.T) { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `BULK INSERT INTO greg-test-01 ( + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `BULK INSERT INTO greg-test-01 ( _id, id_col, string_col, @@ -2351,7 +2368,7 @@ func TestPlanner_SelectSelectSource(t *testing.T) { } t.Run("ParenSource", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b, _id FROM (select * from %j)`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT a, b, _id FROM (select * from %j)`, c)) if err != nil { t.Fatal(err) } @@ -2373,7 +2390,7 @@ func TestPlanner_SelectSelectSource(t *testing.T) { }) t.Run("ParenSourceWithAlias", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT foo.a, b, _id FROM (select * from %j) as foo`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT foo.a, b, _id FROM (select * from %j) as foo`, c)) if err != nil { t.Fatal(err) } @@ -2449,7 +2466,7 @@ func TestPlanner_In(t *testing.T) { t.Run("Count", func(t *testing.T) { t.Skip("Need to add join conditions to get this to pass") - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT %j._id, %j.a, %k._id, %k.parentid, %k.x FROM %j INNER JOIN %k ON %j._id = %k.parentid`, c, c, c, c, c, c, c, c, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT %j._id, %j.a, %k._id, %k.parentid, %k.x FROM %j INNER JOIN %k ON %j._id = %k.parentid`, c, c, c, c, c, c, c, c, c)) // results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT COUNT(*) FROM %j INNER JOIN %k ON %j._id = %k.parentid`, c, c, c, c)) // results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT a FROM %j where a = 20`, c)) // SELECT COUNT(*) FROM %j INNER JOIN %k ON %j._id = %k.parentid if err != nil { @@ -2581,7 +2598,7 @@ func TestPlanner_Distinct(t *testing.T) { } t.Run("SelectDistinct_id", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT distinct _id from %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT distinct _id from %j`, c)) if err != nil { t.Fatal(err) } @@ -2603,7 +2620,7 @@ func TestPlanner_Distinct(t *testing.T) { t.Run("SelectDistinctNonId", func(t *testing.T) { t.Skip("WIP") - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`SELECT distinct parentid from %k`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`SELECT distinct parentid from %k`, c)) if err != nil { t.Fatal(err) } @@ -2624,7 +2641,7 @@ func TestPlanner_Distinct(t *testing.T) { }) t.Run("SelectDistinctMultiple", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`select distinct _id, parentid from %k`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`select distinct _id, parentid from %k`, c)) if err != nil { t.Fatal(err) } @@ -2679,7 +2696,7 @@ func TestPlanner_SelectTop(t *testing.T) { } t.Run("SelectTopStar", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`select top(1) * from %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`select top(1) * from %j`, c)) if err != nil { t.Fatal(err) } @@ -2700,7 +2717,7 @@ func TestPlanner_SelectTop(t *testing.T) { }) t.Run("SelectTopNStar", func(t *testing.T) { - results, columns, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`select topn(1) * from %j`, c)) + results, columns, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`select topn(1) * from %j`, c)) if err != nil { t.Fatal(err) } @@ -2781,12 +2798,12 @@ func TestPlanner_BulkInsert_FB1831(t *testing.T) { c := test.MustRunCluster(t, 3) defer c.Close() - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table iris (_id id, sepallength decimal(2), sepalwidth decimal(2), petallength decimal(2), petalwidth decimal(2), species string cachetype ranked size 1000);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table iris (_id id, sepallength decimal(2), sepalwidth decimal(2), petallength decimal(2), petalwidth decimal(2), species string cachetype ranked size 1000);`) if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into iris (_id, sepallength, sepalwidth, petallength, petalwidth, species) map('id' id, 'sepalLength' DECIMAL(2), @@ -2804,7 +2821,7 @@ func TestPlanner_BulkInsert_FB1831(t *testing.T) { if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(1).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(1).Server, `bulk insert into iris (_id, sepallength, sepalwidth, petallength, petalwidth, species) map('id' id, 'sepalLength' DECIMAL(2), @@ -2822,7 +2839,7 @@ func TestPlanner_BulkInsert_FB1831(t *testing.T) { if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(2).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(2).Server, `bulk insert into iris (_id, sepallength, sepalwidth, petallength, petalwidth, species) map('id' id, 'sepalLength' DECIMAL(2), @@ -2840,7 +2857,7 @@ func TestPlanner_BulkInsert_FB1831(t *testing.T) { if err != nil { t.Fatal(err) } - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `bulk insert into iris (_id, sepallength, sepalwidth, petallength, petalwidth, species) map('id' id, 'sepalLength' DECIMAL(2), @@ -2858,7 +2875,7 @@ func TestPlanner_BulkInsert_FB1831(t *testing.T) { if err != nil { t.Fatal(err) } - results, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id from iris`) + results, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id from iris`) if err != nil { t.Fatal(err) } @@ -2930,7 +2947,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { t.Run("BulkFromLocalFile", func(t *testing.T) { // check that can pull parquet file from local file - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table j1 (_id ID, a INT, b DECIMAL(2), c STRING, d STRINGSET, e IDSET, f BOOL, t TIMESTAMP);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table j1 (_id ID, a INT, b DECIMAL(2), c STRING, d STRINGSET, e IDSET, f BOOL, t TIMESTAMP);`) assert.NoError(t, err) tmpfile, err := os.CreateTemp("", "BulkParquetFile.parquet") assert.NoError(t, err) @@ -2947,7 +2964,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { {Name: "timestampstring", Type: arrow.BinaryTypes.String, Value: []string{"2022-01-28T12:14:04Z", "1970-01-28", "1988-05-30T12:02:00.567999999Z"}}, }) - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j1 (_id,a,b,c,d,e,f,t ) map( 'id' id, @@ -2963,7 +2980,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { WITH FORMAT 'PARQUET' INPUT 'FILE';`, tmpfile.Name())) assert.NoError(t, err) - results, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, a,c from j1`) + results, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, a,c from j1`) assert.NoError(t, err) if diff := cmp.Diff([][]interface{}{ {int64(1), int64(42), "pi"}, @@ -2972,7 +2989,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { }, results); diff != "" { t.Fatal(diff) } - results, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `select b from j1`) + results, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select b from j1`) assert.NoError(t, err) d, _ := pql.FromFloat64WithScale(3.14159, 2) if !pql.Decimal.EqualTo(d, results[0][0].(pql.Decimal)) { @@ -2980,7 +2997,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { } // just getting some opOrderBy coverage here // order by string - results, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, c from j1 order by c`) + results, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, c from j1 order by c`) assert.NoError(t, err) if diff := cmp.Diff([][]interface{}{ {int64(2), "goldenratio"}, @@ -2990,7 +3007,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { t.Fatal(diff) } // order by bool - results, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, f from j1 order by f`) + results, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, f from j1 order by f`) assert.NoError(t, err) if diff := cmp.Diff([][]interface{}{ {int64(2), false}, @@ -3000,7 +3017,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { t.Fatal(diff) } // order by timestamp - results, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, t from j1 order by t`) + results, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, t from j1 order by t`) assert.NoError(t, err) t2, _ := time.ParseInLocation(time.RFC3339Nano, "1970-01-28T00:00:00Z", time.UTC) t3, _ := time.ParseInLocation(time.RFC3339Nano, "1988-05-30T12:02:00Z", time.UTC) @@ -3018,7 +3035,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { t.Run("BulkParquetFromUrl", func(t *testing.T) { // check that can pull parquet file from URL and load - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table j2 (_id ID, a INT, b STRING);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table j2 (_id ID, a INT, b STRING);`) assert.NoError(t, err) tmpfile, err := os.CreateTemp("", "BulkParquetFile.parquet") assert.NoError(t, err) @@ -3041,7 +3058,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { w.Write(payload) }) - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j2 (_id,a,b ) map( 'id' id, @@ -3052,7 +3069,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { WITH FORMAT 'PARQUET' INPUT 'URL';`, ts.URL+"/static")) assert.NoError(t, err) - results, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, a,b from j2`) + results, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, a,b from j2`) assert.NoError(t, err) if diff := cmp.Diff([][]interface{}{ {int64(1), int64(42), "pi"}, @@ -3062,7 +3079,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { }) t.Run("FileTimeStamp", func(t *testing.T) { - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table continuum (_id ID, created timestamp, updated timestamp);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table continuum (_id ID, created timestamp, updated timestamp);`) assert.NoError(t, err) tmpfile, err := os.CreateTemp("", "BulkParquetFile.parquet") assert.NoError(t, err) @@ -3075,7 +3092,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { {Name: "stringtime", Type: arrow.BinaryTypes.String, Value: []string{now.Format(time.RFC3339)}}, }) - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into continuum (_id,created,updated ) map( 'id' id, @@ -3086,7 +3103,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { WITH FORMAT 'PARQUET' INPUT 'FILE';`, tmpfile.Name())) assert.NoError(t, err) - results, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, created,updated from continuum`) + results, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, created,updated from continuum`) assert.NoError(t, err) row := results[0] assert.Equal(t, row[1], row[2]) @@ -3095,7 +3112,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { t.Run("BulkCSVFromUrl", func(t *testing.T) { // check that can pull parquet file from URL and load - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table j3 (_id ID, a INT, b STRING);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table j3 (_id ID, a INT, b STRING);`) assert.NoError(t, err) tmpfile, err := os.CreateTemp("", "BulkCSVFile.csv") assert.NoError(t, err) @@ -3119,7 +3136,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { w.Write(payload) }) - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j3 (_id,a,b ) map( 0 id, @@ -3131,7 +3148,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { BATCHSIZE 60 INPUT 'URL';`, ts.URL+"/static")) assert.NoError(t, err) - results, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, a,b from j3`) + results, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, a,b from j3`) assert.NoError(t, err) if diff := cmp.Diff(expect, results); diff != "" { t.Fatal(diff) @@ -3140,7 +3157,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { t.Run("BulkNDJSONFromUrl", func(t *testing.T) { // check that can pull parquet file from URL and load - _, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `create table j4 (_id ID, a INT, b STRING, c TIMESTAMP);`) + _, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `create table j4 (_id ID, a INT, b STRING, c TIMESTAMP);`) assert.NoError(t, err) tmpfile, err := os.CreateTemp("", "BulkJSONFile.json") assert.NoError(t, err) @@ -3159,7 +3176,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { w.Write(payload) }) - _, _, _, err = sql_test.MustQueryRows(t, c.GetNode(0).Server, fmt.Sprintf(`bulk insert + _, _, _, err = sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, fmt.Sprintf(`bulk insert into j4 (_id,a,b,c ) map( '$.id' id, @@ -3171,7 +3188,7 @@ func TestPlanner_BulkInsertParquet(t *testing.T) { WITH FORMAT 'NDJSON' INPUT 'URL';`, ts.URL+"/static")) assert.NoError(t, err) - results, _, _, err := sql_test.MustQueryRows(t, c.GetNode(0).Server, `select _id, a,b from j4`) + results, _, _, err := sql_test.MustQueryRows(t, nil, c.GetNode(0).Server, `select _id, a,b from j4`) assert.NoError(t, err) if diff := cmp.Diff([][]interface{}{ {int64(1), int64(42), "pi"}, @@ -3186,7 +3203,7 @@ func TestPlanner_BulkInsert_FP1916(t *testing.T) { defer c.Close() // 8924809397503602651 is larger than 2^53 which is the largest integer value representable in float64 node := c.GetNode(0).Server - _, _, _, err := sql_test.MustQueryRows(t, node, `create table greg-test ( + _, _, _, err := sql_test.MustQueryRows(t, nil, node, `create table greg-test ( _id STRING, id_col ID, string_col STRING cachetype ranked size 1000, @@ -3199,7 +3216,7 @@ func TestPlanner_BulkInsert_FP1916(t *testing.T) { );`) assert.NoError(t, err) - _, _, _, err = sql_test.MustQueryRows(t, node, `BULK INSERT INTO greg-test ( + _, _, _, err = sql_test.MustQueryRows(t, nil, node, `BULK INSERT INTO greg-test ( _id, id_col, string_col, @@ -3234,7 +3251,7 @@ func TestPlanner_BulkInsert_FP1916(t *testing.T) { format 'CSV' input 'STREAM';`) assert.NoError(t, err) - results, _, _, err := sql_test.MustQueryRows(t, node, `select int_col from greg-test`) + results, _, _, err := sql_test.MustQueryRows(t, nil, node, `select int_col from greg-test`) assert.NoError(t, err) got := results[0][0].(int64) expected := int64(8924809397503602651) @@ -3246,12 +3263,12 @@ func TestPlanner_BulkInsert_FP1915(t *testing.T) { defer c.Close() // 8924809397503602651 is larger than 2^53 which is the largest integer value representable in float64 node := c.GetNode(0).Server - _, _, _, err := sql_test.MustQueryRows(t, node, `create table ids ( + _, _, _, err := sql_test.MustQueryRows(t, nil, node, `create table ids ( _id id, a int, b int);`) assert.NoError(t, err) - _, _, _, err = sql_test.MustQueryRows(t, node, `BULK INSERT INTO ids (_id, a, b) + _, _, _, err = sql_test.MustQueryRows(t, nil, node, `BULK INSERT INTO ids (_id, a, b) map ('$._id' id, '$.a' int, '$.b' int) from x'{ "_id":8924809397503602651 , "a": 10, "b": 20 } { "_id":"8924809397503602652" , "a": 10, "b": 20 }' @@ -3259,7 +3276,7 @@ WITH FORMAT 'NDJSON' INPUT 'STREAM';`) assert.NoError(t, err) - results, _, _, err := sql_test.MustQueryRows(t, node, `select _id from ids`) + results, _, _, err := sql_test.MustQueryRows(t, nil, node, `select _id from ids`) assert.NoError(t, err) got := make([]int64, 0) for i := range results { @@ -3277,14 +3294,14 @@ WITH } // FB_2062 - _, _, _, err = sql_test.MustQueryRows(t, node, `create table sup305-fails (_id id, bucket string, value int);`) + _, _, _, err = sql_test.MustQueryRows(t, nil, node, `create table sup305-fails (_id id, bucket string, value int);`) assert.NoError(t, err) - _, _, _, err = sql_test.MustQueryRows(t, node, ` + _, _, _, err = sql_test.MustQueryRows(t, nil, node, ` insert into sup305-fails values (1, 'a', 1000), (2, 'b', 1000), (3, 'c', 1000), (4, 'c', 1000), (5, 'c', 1000), (6, 'c', 1000), (7, 'c', 1000), (8, 'a', 1000), (9, 'b', 1000), (10, 'c', 1000), (11, 'c', 1000), (12, 'c', 1000), (13, 'c', 1000), (14, 'c', 1000), (15, 'a', 1000), (16, 'b', 1000), (17, 'c', 1000), (18, 'c', 1000), (19, 'c', 1000), (20, 'c', 1000), (21, 'c', 1000);`) assert.NoError(t, err) - results, _, _, err = sql_test.MustQueryRows(t, node, `select bucket, count(*) as cnt from sup305-fails group by bucket having count(*) > 1 order by cnt;`) + results, _, _, err = sql_test.MustQueryRows(t, nil, node, `select bucket, count(*) as cnt from sup305-fails group by bucket having count(*) > 1 order by cnt;`) assert.NoError(t, err) got = make([]int64, 0) for i := range results { diff --git a/sql3/sql_test.go b/sql3/sql_test.go index 6a10be5c8..2ed0fd891 100644 --- a/sql3/sql_test.go +++ b/sql3/sql_test.go @@ -46,13 +46,13 @@ func TestSQL_Execute(t *testing.T) { // Create a table with all field types. if test.HasTable() { - _, _, _, err := sql_test.MustQueryRows(t, svr, test.CreateTable()) + _, _, _, err := sql_test.MustQueryRows(t, nil, svr, test.CreateTable()) assert.NoError(t, err) } if test.HasTable() && test.HasData() { // Populate fields with data. - _, _, _, err := sql_test.MustQueryRows(t, svr, test.InsertInto(t)) + _, _, _, err := sql_test.MustQueryRows(t, nil, svr, test.InsertInto(t)) assert.NoError(t, err) } @@ -61,7 +61,7 @@ func TestSQL_Execute(t *testing.T) { for _, sql := range sqltest.SQLs { t.Run(fmt.Sprintf("sql-%s", sql), func(t *testing.T) { log.Printf("SQL: %s", sql) - rows, headers, plan, err := sql_test.MustQueryRows(t, svr, sql) + rows, headers, plan, err := sql_test.MustQueryRows(t, nil, svr, sql) // Check expected error instead of results. if sqltest.ExpErr != "" { diff --git a/sql3/test/defs/defs.go b/sql3/test/defs/defs.go index 6b309737c..5412758db 100644 --- a/sql3/test/defs/defs.go +++ b/sql3/test/defs/defs.go @@ -39,7 +39,7 @@ var TableTests []TableTest = []TableTest{ subqueryTests, viewTests, - topTests, + topLimitTests, deleteTests, @@ -179,6 +179,8 @@ var TableTests []TableTest = []TableTest{ avgTests, percentileTests, minmaxTests, + corrTests, + varTests, // groupby tests groupByTests, @@ -197,6 +199,10 @@ var TableTests []TableTest = []TableTest{ bulkInsertTable, bulkInsert, + // copy + copyTable, + copyTests, + // bool (batch logic) boolTests, diff --git a/sql3/test/defs/defs_aggregate.go b/sql3/test/defs/defs_aggregate.go index a52698685..0532d3c40 100644 --- a/sql3/test/defs/defs_aggregate.go +++ b/sql3/test/defs/defs_aggregate.go @@ -737,3 +737,178 @@ var minmaxTests = TableTest{ }, }, } + +var corrTests = TableTest{ + Table: tbl( + "corr_test", + srcHdrs( + srcHdr("_id", fldTypeID), + srcHdr("i1", fldTypeInt, "min 0", "max 1000"), + srcHdr("d1", fldTypeDecimal2), + srcHdr("s1", fldTypeString), + srcHdr("id1", fldTypeID), + ), + srcRows( + srcRow(int64(1), int64(10), float64(10), string("foo"), int64(10)), + srcRow(int64(2), int64(10), float64(10), string("foo2"), int64(11)), + srcRow(int64(3), int64(11), float64(11), string("foo34"), int64(12)), + srcRow(int64(4), int64(12), float64(12), string("foo45"), int64(13)), + srcRow(int64(5), int64(12), float64(12), string("foo63"), int64(14)), + srcRow(int64(6), int64(13), float64(13), string("foo22222"), int64(15)), + ), + ), + SQLTests: []SQLTest{ + { + SQLs: sqls( + "SELECT corr(*, i1) AS corr1 FROM corr_test", + ), + ExpErr: "expected right paren, found ','", + }, + { + SQLs: sqls( + "SELECT corr(i1, d1) AS corr1 FROM corr_test", + ), + ExpHdrs: hdrs( + hdr("corr1", featurebase.WireQueryField{ + Type: dax.BaseTypeDecimal + "(6)", + BaseType: dax.BaseTypeDecimal, + TypeInfo: map[string]interface{}{"scale": int64(6)}, + }), + ), + ExpRows: rows( + row(pql.NewDecimal(1000000, 6)), + ), + Compare: CompareExactUnordered, + }, + { + SQLs: sqls( + "SELECT corr(_id, i1) AS corr1 FROM corr_test", + ), + ExpErr: "_id column cannot be used in aggregate function 'corr'", + }, + { + SQLs: sqls( + "SELECT corr(i1) AS corr1 FROM corr_test", + ), + ExpErr: "count of formal parameters (2) does not match count of actual parameters (1)", + }, + { + SQLs: sqls( + "SELECT corr(s1, i1) AS avg_rows FROM corr_test", + ), + ExpErr: "integer, decimal or timestamp expression expected", + }, + { + SQLs: sqls( + "SELECT corr(len(s1), i1) AS corr1 FROM corr_test", + ), + ExpHdrs: hdrs( + hdr("corr1", featurebase.WireQueryField{ + Type: dax.BaseTypeDecimal + "(6)", + BaseType: dax.BaseTypeDecimal, + TypeInfo: map[string]interface{}{"scale": int64(6)}, + }), + ), + ExpRows: rows( + row(pql.NewDecimal(888234, 6)), + ), + Compare: CompareExactUnordered, + }, + }, +} + +var varTests = TableTest{ + Table: tbl( + "var_test", + srcHdrs( + srcHdr("_id", fldTypeID), + srcHdr("i1", fldTypeInt, "min 0", "max 1000"), + srcHdr("d1", fldTypeDecimal2), + srcHdr("s1", fldTypeString), + srcHdr("id1", fldTypeID), + ), + srcRows( + srcRow(int64(1), int64(10), float64(10), string("foo"), int64(10)), + srcRow(int64(2), int64(10), float64(10), string("foo"), int64(11)), + srcRow(int64(3), int64(11), float64(11), string("foo"), int64(12)), + srcRow(int64(4), int64(12), float64(12), string("foo"), int64(13)), + srcRow(int64(5), int64(12), float64(12), string("foo"), int64(14)), + srcRow(int64(6), int64(13), float64(13), string("foo"), int64(15)), + ), + ), + SQLTests: []SQLTest{ + { + SQLs: sqls( + "SELECT var(*) AS var1 FROM var_test", + ), + ExpErr: "column reference expected", + }, + { + SQLs: sqls( + "SELECT var(_id) AS var1 FROM var_test", + ), + ExpErr: "_id column cannot be used in aggregate function 'var'", + }, + { + SQLs: sqls( + "SELECT var(i1, d1) AS var1 FROM var_test", + ), + ExpErr: "count of formal parameters (1) does not match count of actual parameters (2)", + }, + { + SQLs: sqls( + "SELECT var(s1) AS var1 FROM var_test", + ), + ExpErr: "integer, decimal or timestamp expression expected", + }, + { + SQLs: sqls( + "SELECT var(id1) AS var1 FROM var_test", + ), + ExpHdrs: hdrs( + hdr("var1", featurebase.WireQueryField{ + Type: dax.BaseTypeDecimal + "(6)", + BaseType: dax.BaseTypeDecimal, + TypeInfo: map[string]interface{}{"scale": int64(6)}, + }), + ), + ExpRows: rows( + row(pql.NewDecimal(2916666, 6)), + ), + Compare: CompareExactUnordered, + }, + { + SQLs: sqls( + "SELECT var(i1) AS var1 FROM var_test", + "SELECT var(d1) AS var1 FROM var_test", + ), + ExpHdrs: hdrs( + hdr("var1", featurebase.WireQueryField{ + Type: dax.BaseTypeDecimal + "(6)", + BaseType: dax.BaseTypeDecimal, + TypeInfo: map[string]interface{}{"scale": int64(6)}, + }), + ), + ExpRows: rows( + row(pql.NewDecimal(1222222, 6)), + ), + Compare: CompareExactUnordered, + }, + { + SQLs: sqls( + "SELECT var(len(s1)) AS var1 FROM var_test", + ), + ExpHdrs: hdrs( + hdr("var1", featurebase.WireQueryField{ + Type: dax.BaseTypeDecimal + "(6)", + BaseType: dax.BaseTypeDecimal, + TypeInfo: map[string]interface{}{"scale": int64(6)}, + }), + ), + ExpRows: rows( + row(pql.NewDecimal(0, 6)), + ), + Compare: CompareExactUnordered, + }, + }, +} diff --git a/sql3/test/defs/defs_bulkinsert.go b/sql3/test/defs/defs_bulkinsert.go index 57f0b4036..3fcafa852 100644 --- a/sql3/test/defs/defs_bulkinsert.go +++ b/sql3/test/defs/defs_bulkinsert.go @@ -1,6 +1,6 @@ package defs -// join tests +// bulk insert var bulkInsertTable = TableTest{ name: "bulkInsertTable", Table: tbl( diff --git a/sql3/test/defs/defs_copy.go b/sql3/test/defs/defs_copy.go new file mode 100644 index 000000000..3fb3be48d --- /dev/null +++ b/sql3/test/defs/defs_copy.go @@ -0,0 +1,58 @@ +// Copyright 2023 Molecula Corp. All rights reserved. + +package defs + +// copy +var copyTable = TableTest{ + name: "copyTable", + Table: tbl( + "copytest", + srcHdrs( + srcHdr("_id", fldTypeID), + srcHdr("id_col", fldTypeID), + srcHdr("string_col", fldTypeString), + srcHdr("int_col", fldTypeInt), + srcHdr("decimal_col", fldTypeDecimal2), + srcHdr("bool_col", fldTypeBool), + srcHdr("time_col", fldTypeTimestamp), + srcHdr("stringset_col", fldTypeStringSet), + srcHdr("idset_col", fldTypeIDSet), + ), + srcRows( + srcRow(int64(1), int64(10), string("foo"), int64(10), float64(10), bool(false), knownTimestamp(), []string{"foo", "bar"}, []int64{1, 2}), + srcRow(int64(2), int64(11), string("foo1"), int64(11), float64(11), bool(true), knownTimestamp(), []string{"foo1", "bar1"}, []int64{11, 21}), + srcRow(int64(3), int64(12), string("foo2"), int64(12), float64(12), bool(false), knownTimestamp(), []string{"foo2", "bar2"}, []int64{12, 22}), + srcRow(int64(4), int64(13), string("foo3"), int64(13), float64(13), bool(true), knownTimestamp(), []string{"foo3", "bar3"}, []int64{13, 23}), + ), + ), + SQLTests: nil, +} + +var copyTests = TableTest{ + name: "copyTests", + SQLTests: []SQLTest{ + { + name: "copy-no-table-to-table", + SQLs: sqls( + `copy foo to bar;`, + ), + ExpErr: "table or view 'foo' not found", + }, + { + name: "copy-table-to-table-same-name", + SQLs: sqls( + `copy copytest to copytest;`, + ), + ExpErr: "already exists", + }, + { + name: "copy-table-to-table", + SQLs: sqls( + `copy copytest to copytesttwo;`, + ), + ExpHdrs: hdrs(), + ExpRows: rows(), + Compare: CompareExactOrdered, + }, + }, +} diff --git a/sql3/test/defs/defs_create_table.go b/sql3/test/defs/defs_create_table.go index 724dd03f2..b80ffe81f 100644 --- a/sql3/test/defs/defs_create_table.go +++ b/sql3/test/defs/defs_create_table.go @@ -31,6 +31,13 @@ var createTable = TableTest{ ), ExpErr: "expected literal, found bad", }, + { + name: "minAboveMax", + SQLs: sqls( + "create table bar (_id id, i1 int min 20 max 19)", + ), + ExpErr: "int field min cannot be greater than max", + }, { name: "commentString", SQLs: sqls( diff --git a/sql3/test/defs/defs_delete.go b/sql3/test/defs/defs_delete.go index 68d2ff851..cfa6b5d49 100644 --- a/sql3/test/defs/defs_delete.go +++ b/sql3/test/defs/defs_delete.go @@ -118,6 +118,30 @@ var deleteTests = TableTest{ ExpRows: rows(), Compare: CompareExactUnordered, }, + // inner joins in delete are apparently not supported yet + // when they are, this will fill in test coverage for expressionanalyzer.go + /* + { + SQLs: sqls( + "delete from del_all_types a1 inner join del_all_types a2 on a1._id=a2._id where a1._id in (select _id from del_all_types);", + ), + ExpHdrs: hdrs( + hdr("_id", fldTypeID), + ), + ExpRows: rows(), + Compare: CompareExactUnordered, + }, + { + SQLs: sqls( + "select _id from del_all_types;", + ), + ExpHdrs: hdrs( + hdr("_id", fldTypeID), + ), + ExpRows: rows(), + Compare: CompareExactUnordered, + }, + */ // dates { diff --git a/sql3/test/defs/defs_in.go b/sql3/test/defs/defs_in.go index e2d0dfaf1..a039f75d4 100644 --- a/sql3/test/defs/defs_in.go +++ b/sql3/test/defs/defs_in.go @@ -128,6 +128,51 @@ var inTests = TableTest{ ), Compare: CompareExactUnordered, }, + { + SQLs: sqls( + "select t1._id in (select t2._id from in_all_types as t2) from in_all_types as t1", + ), + ExpHdrs: hdrs( + hdr("", fldTypeBool), + ), + ExpRows: rows( + row(bool(true)), + ), + Compare: CompareExactUnordered, + }, + // This test fails - re-enable it when fixing the bug + /* + { + SQLs: sqls( + "select t1._id in (select t2.id1 from in_all_types as t2) from in_all_types as t1", + ), + ExpHdrs: hdrs( + hdr("", fldTypeBool), + ), + ExpRows: rows( + row(bool(false)), + ), + Compare: CompareExactUnordered, + }, + */ + // This test also fails. + // Once re-enabled and passing, it provides coverage for expressionanalyzer.go + // (*ExecutionPlanner).analyzeBinaryExpression in the case where the left hand side + // of the IN contains a JOIN. + /* + { + SQLs: sqls( + "select a1._id from in_all_types a1 inner join in_all_types a2 on a1._id=a2._id where a1._id in (select _id from in_all_types);", + ), + ExpHdrs: hdrs( + hdr("", fldTypeBool), + ), + ExpRows: rows( + row(bool(true)), + ), + Compare: CompareExactUnordered, + }, + */ }, } @@ -259,5 +304,63 @@ var notInTests = TableTest{ ), Compare: CompareExactUnordered, }, + // this test currently causes a panic + /* + { + SQLs: sqls( + "select t1._id in (select t2.s1 from in_all_types as t2) from in_all_types as t1", + ), + ExpHdrs: hdrs( + hdr("", fldTypeBool), + ), + ExpErr: "types 'id' and 'string' are not equatable", + Compare: CompareExactUnordered, + }, + */ + { + SQLs: sqls( + "select t1._id not in (select t2._id from in_all_types as t2) from in_all_types as t1", + ), + ExpHdrs: hdrs( + hdr("", fldTypeBool), + ), + ExpRows: rows( + row(bool(false)), + ), + Compare: CompareExactUnordered, + }, + // This test fails - re-enable it when the bug is fixed + /* + { + SQLs: sqls( + "select t1._id not in (select t2.id1 from in_all_types as t2) from in_all_types as t1", + ), + ExpHdrs: hdrs( + hdr("", fldTypeBool), + ), + ExpRows: rows( + row(bool(true)), + ), + Compare: CompareExactUnordered, + }, + */ + // This test also fails. + // Once re-enabled and passing, it provides coverage for expressionanalyzer.go + // (*ExecutionPlanner).analyzeBinaryExpression in the case where the left hand side + // of the IN contains a JOIN. + /* + { + SQLs: sqls( + "select a1._id from in_all_types a1 inner join in_all_types a2 on a1._id=a2._id where a1._id not in (select _id from in_all_types);", + ), + ExpHdrs: hdrs( + hdr("", fldTypeBool), + ), + ExpRows: rows( + row(bool(true)), + ), + Compare: CompareExactUnordered, + }, + */ }, } diff --git a/sql3/test/defs/defs_top.go b/sql3/test/defs/defs_top.go index f4c61e41a..324cc9c00 100644 --- a/sql3/test/defs/defs_top.go +++ b/sql3/test/defs/defs_top.go @@ -1,7 +1,7 @@ package defs -var topTests = TableTest{ - name: "top-tests", +var topLimitTests = TableTest{ + name: "top-limit-tests", Table: tbl( "skills", srcHdrs( @@ -20,7 +20,25 @@ var topTests = TableTest{ SQLTests: []SQLTest{ { SQLs: sqls( - "select top(1) * from skills where setcontains(skills, 'Marketing Manager');", + "select top(1) * from skills where setcontains(skills, 'Marketing Manager');", + ), + ExpHdrs: hdrs( + hdr("_id", fldTypeID), + hdr("bools", fldTypeStringSet), + hdr("bools-exist", fldTypeStringSet), + hdr("id1", fldTypeID), + hdr("skills", fldTypeStringSet), + hdr("titles", fldTypeStringSet), + ), + ExpRows: rows( + row(int64(1), nil, []string{"available_for_hire"}, int64(288), []string{"Marketing Manager"}, []string{"Alumni Relations", "OEM negotiations"}), + ), + Compare: CompareExactUnordered, + SortStringKeys: true, + }, + { + SQLs: sqls( + "select * from skills where setcontains(skills, 'Marketing Manager') limit 1;", ), ExpHdrs: hdrs( hdr("_id", fldTypeID), @@ -55,6 +73,21 @@ var topTests = TableTest{ Compare: CompareExactUnordered, SortStringKeys: true, }, + { + SQLs: sqls( + "select count(*), skills from skills group by skills limit 10;", + ), + ExpHdrs: hdrs( + hdr("", fldTypeInt), + hdr("skills", fldTypeStringSet), + ), + ExpRows: rows( + row(int64(1), string("Marketing Manager")), + row(int64(1), string("Software Engineer I")), + ), + Compare: CompareExactUnordered, + SortStringKeys: true, + }, { SQLs: sqls( "select top(1) count(*) from skills;", @@ -68,5 +101,24 @@ var topTests = TableTest{ Compare: CompareExactUnordered, SortStringKeys: true, }, + { + SQLs: sqls( + "select count(*) from skills limit 1;", + ), + ExpHdrs: hdrs( + hdr("", fldTypeInt), + ), + ExpRows: rows( + row(int64(2)), + ), + Compare: CompareExactUnordered, + SortStringKeys: true, + }, + { + SQLs: sqls( + "select top(1) count(*) from skills limit 1;", + ), + ExpErr: "TOP and LIMIT cannot cannot be used at the same time", + }, }, } diff --git a/sql3/test/helpers.go b/sql3/test/helpers.go index 9e7b62389..bcd26fc22 100644 --- a/sql3/test/helpers.go +++ b/sql3/test/helpers.go @@ -5,6 +5,7 @@ import ( "context" "encoding/json" "testing" + "time" featurebase "github.com/featurebasedb/featurebase/v3" fbcontext "github.com/featurebasedb/featurebase/v3/context" @@ -14,14 +15,26 @@ import ( ) // MustQueryRows returns the row results as a slice of []interface{}, along with the columns, the query plan as a []byte or an error. -func MustQueryRows(tb testing.TB, svr *featurebase.Server, q string) ([][]interface{}, []*featurebase.WireQueryField, []byte, error) { +func MustQueryRows(tb testing.TB, c context.Context, svr *featurebase.Server, q string) ([][]interface{}, []*featurebase.WireQueryField, []byte, error) { tb.Helper() requestId, err := uuid.NewV4() if err != nil { return nil, nil, nil, err } - ctx := fbcontext.WithRequestID(context.Background(), requestId.String()) + // Originally MustQueryRows just created a context for itself do test with. + // However, for some tests, we may want access to the test's context so that, + // for example, we can cancel the context mid-test and make sure that gets + // handled correctly. Since we need to be able to cancel the context before + // the query finishes, if a context is set we also introduce a delay. + var ctx context.Context + delay := 0 * time.Millisecond + if c == nil { + ctx = fbcontext.WithRequestID(context.Background(), requestId.String()) + } else { + ctx = fbcontext.WithRequestID(c, requestId.String()) + delay = 100 * time.Millisecond + } stmt, err := svr.CompileExecutionPlan(ctx, q) if err != nil { @@ -43,10 +56,17 @@ func MustQueryRows(tb testing.TB, svr *featurebase.Server, q string) ([][]interf } results := make([][]interface{}, 0) + // figuring out where to put the sleep took some trial and error. + // Too early and the context gets cancelled before the query can + // start running, too late and you end up with the query getting + // cancelled instead of the context. + time.Sleep(delay) + next, err := rowIter.Next(ctx) if err != nil && err != plannertypes.ErrNoMoreRows { return nil, nil, nil, err } + for err != plannertypes.ErrNoMoreRows { result := make([]interface{}, len(ocolumns)) for i := range result {