Merge branch 'master' into cloud-1475

This commit is contained in:
David Kagan 2023-04-06 12:30:04 -04:00 committed by GitHub
commit cd13416614
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
71 changed files with 4913 additions and 1173 deletions

View file

@ -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) {

View file

@ -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)
}

View file

@ -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)

View file

@ -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)

View file

@ -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 {

View file

@ -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)

View file

@ -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")

View file

@ -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
}

View file

@ -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

View file

@ -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)
}

View file

@ -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()

View file

@ -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)
}

View file

@ -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")

View file

@ -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))

View file

@ -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
}

View file

@ -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
}

View file

@ -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 {

View file

@ -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))

40
dax/controller/worker.go Normal file
View file

@ -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
}

View file

@ -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")

View file

@ -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
}

View file

@ -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)
}

View file

@ -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

View file

@ -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

View file

@ -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) {

View file

@ -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
}

2
go.mod
View file

@ -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

3
go.sum
View file

@ -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=

View file

@ -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),
)
}

View file

@ -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()
}

View file

@ -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)}},

View file

@ -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 {

View file

@ -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'`)

View file

@ -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",

120
sql3/planner/compilecopy.go Normal file
View file

@ -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
}

View file

@ -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
}

View file

@ -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)
}
}

View file

@ -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
}

View file

@ -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
}

View file

@ -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

View file

@ -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)
}

View file

@ -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 &timestampLiteralPlanExpression{
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
}
}

View file

@ -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")

View file

@ -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{}

View file

@ -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 {

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)

515
sql3/planner/opcopy.go Normal file
View file

@ -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 &copyIterator{
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
}

View file

@ -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
}

View file

@ -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
}

View file

@ -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)
}

102
sql3/planner/opdropmodel.go Normal file
View file

@ -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
}

253
sql3/planner/oppredict.go Normal file
View file

@ -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
}

View file

@ -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),

View file

@ -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:

View file

@ -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
}

View file

@ -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")
}

File diff suppressed because it is too large Load diff

View file

@ -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 != "" {

View file

@ -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,

View file

@ -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,
},
},
}

View file

@ -1,6 +1,6 @@
package defs
// join tests
// bulk insert
var bulkInsertTable = TableTest{
name: "bulkInsertTable",
Table: tbl(

View file

@ -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,
},
},
}

View file

@ -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(

View file

@ -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
{

View file

@ -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,
},
*/
},
}

View file

@ -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",
},
},
}

View file

@ -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 {