Storage: Watch support (#82282)

* initial naive implementation

* Update pkg/services/store/entity/sqlstash/sql_storage_server.go

Co-authored-by: Igor Suleymanov <radiohead@users.noreply.github.com>

* tidy up

* add action column, batch watch events

* initial implementation of broadcast-based watcher

* fix up watch init

* remove batching, it just adds needless complexity

* use StreamWatcher

* make broadcaster generic

* add circular buffer to replay recent events to new watchers

* loop within poll until all events are read

* add index on entity_history.resource_version to support poller

* increment r.Since when we send events to consumer

* switch broadcaster and cache to use channels instead of mutexes

* cleanup

---------

Co-authored-by: Igor Suleymanov <radiohead@users.noreply.github.com>
This commit is contained in:
Dan Cech
2024-03-05 10:14:38 -05:00
committed by GitHub
co-authored by Igor Suleymanov
parent 71fe675fb7
commit 7b4925ea37
10 changed files with 1293 additions and 394 deletions
@@ -0,0 +1,254 @@
package sqlstash
import (
"context"
"fmt"
)
type ConnectFunc[T any] func(chan T) error
type Broadcaster[T any] interface {
Subscribe(context.Context) (<-chan T, error)
Unsubscribe(chan T)
}
func NewBroadcaster[T any](ctx context.Context, connect ConnectFunc[T]) (Broadcaster[T], error) {
b := &broadcaster[T]{}
err := b.start(ctx, connect)
if err != nil {
return nil, err
}
return b, nil
}
type broadcaster[T any] struct {
running bool
ctx context.Context
subs map[chan T]struct{}
cache Cache[T]
subscribe chan chan T
unsubscribe chan chan T
}
func (b *broadcaster[T]) Subscribe(ctx context.Context) (<-chan T, error) {
if !b.running {
return nil, fmt.Errorf("broadcaster not running")
}
sub := make(chan T, 100)
b.subscribe <- sub
go func() {
<-ctx.Done()
b.unsubscribe <- sub
}()
return sub, nil
}
func (b *broadcaster[T]) Unsubscribe(sub chan T) {
b.unsubscribe <- sub
}
func (b *broadcaster[T]) start(ctx context.Context, connect ConnectFunc[T]) error {
if b.running {
return fmt.Errorf("broadcaster already running")
}
stream := make(chan T, 100)
err := connect(stream)
if err != nil {
return err
}
b.ctx = ctx
b.cache = NewCache[T](ctx, 100)
b.subscribe = make(chan chan T, 100)
b.unsubscribe = make(chan chan T, 100)
b.subs = make(map[chan T]struct{})
go b.stream(stream)
b.running = true
return nil
}
func (b *broadcaster[T]) stream(input chan T) {
for {
select {
// context cancelled
case <-b.ctx.Done():
close(input)
for sub := range b.subs {
close(sub)
delete(b.subs, sub)
}
b.running = false
return
// new subscriber
case sub := <-b.subscribe:
// send initial batch of cached items
err := b.cache.ReadInto(sub)
if err != nil {
close(sub)
continue
}
b.subs[sub] = struct{}{}
// unsubscribe
case sub := <-b.unsubscribe:
if _, ok := b.subs[sub]; ok {
close(sub)
delete(b.subs, sub)
}
// read item from input
case item, ok := <-input:
// input closed, drain subscribers and exit
if !ok {
for sub := range b.subs {
close(sub)
delete(b.subs, sub)
}
b.running = false
return
}
b.cache.Add(item)
for sub := range b.subs {
select {
case sub <- item:
default:
// Slow consumer, drop
b.unsubscribe <- sub
}
}
}
}
}
const DefaultCacheSize = 100
type Cache[T any] interface {
Len() int
Add(item T)
Get(i int) T
Range(f func(T) error) error
Slice() []T
ReadInto(dst chan T) error
}
type cache[T any] struct {
cache []T
size int
cacheZero int
cacheLen int
add chan T
read chan chan T
ctx context.Context
}
func NewCache[T any](ctx context.Context, size int) Cache[T] {
c := &cache[T]{}
c.ctx = ctx
if size <= 0 {
size = DefaultCacheSize
}
c.size = size
c.cache = make([]T, c.size)
c.add = make(chan T)
c.read = make(chan chan T)
go c.run()
return c
}
func (c *cache[T]) Len() int {
return c.cacheLen
}
func (c *cache[T]) Add(item T) {
c.add <- item
}
func (c *cache[T]) run() {
for {
select {
case <-c.ctx.Done():
return
case item := <-c.add:
i := (c.cacheZero + c.cacheLen) % len(c.cache)
c.cache[i] = item
if c.cacheLen < len(c.cache) {
c.cacheLen++
} else {
c.cacheZero = (c.cacheZero + 1) % len(c.cache)
}
case r := <-c.read:
read:
for i := 0; i < c.cacheLen; i++ {
select {
case r <- c.cache[(c.cacheZero+i)%len(c.cache)]:
// don't wait for slow consumers
default:
break read
}
}
close(r)
}
}
}
func (c *cache[T]) Get(i int) T {
r := make(chan T, c.size)
c.read <- r
idx := 0
for item := range r {
if idx == i {
return item
}
idx++
}
var zero T
return zero
}
func (c *cache[T]) Range(f func(T) error) error {
r := make(chan T, c.size)
c.read <- r
for item := range r {
err := f(item)
if err != nil {
return err
}
}
return nil
}
func (c *cache[T]) Slice() []T {
s := make([]T, 0, c.size)
r := make(chan T, c.size)
c.read <- r
for item := range r {
s = append(s, item)
}
return s
}
func (c *cache[T]) ReadInto(dst chan T) error {
r := make(chan T, c.size)
c.read <- r
for item := range r {
select {
case dst <- item:
default:
return fmt.Errorf("slow consumer")
}
}
return nil
}
@@ -0,0 +1,106 @@
package sqlstash
import (
"context"
"testing"
"github.com/stretchr/testify/require"
)
func TestCache(t *testing.T) {
c := NewCache[int](context.Background(), 10)
e := []int{}
err := c.Range(func(i int) error {
e = append(e, i)
return nil
})
require.Nil(t, err)
require.Equal(t, 0, len(e))
c.Add(1)
e = []int{}
err = c.Range(func(i int) error {
e = append(e, i)
return nil
})
require.Nil(t, err)
require.Equal(t, 1, len(e))
require.Equal(t, []int{1}, e)
require.Equal(t, 1, c.Get(0))
c.Add(2)
c.Add(3)
c.Add(4)
c.Add(5)
c.Add(6)
// should be able to range over values
e = []int{}
err = c.Range(func(i int) error {
e = append(e, i)
return nil
})
require.Nil(t, err)
require.Equal(t, 6, len(e))
require.Equal(t, []int{1, 2, 3, 4, 5, 6}, e)
// should be able to get length
require.Equal(t, 6, c.Len())
// should be able to get values
require.Equal(t, 1, c.Get(0))
require.Equal(t, 6, c.Get(5))
// zero value beyond cache size
require.Equal(t, 0, c.Get(6))
require.Equal(t, 0, c.Get(20))
require.Equal(t, 0, c.Get(-10))
// slice should return all values
require.Equal(t, []int{1, 2, 3, 4, 5, 6}, c.Slice())
c.Add(7)
c.Add(8)
c.Add(9)
c.Add(10)
c.Add(11)
// should be able to range over values
e = []int{}
err = c.Range(func(i int) error {
e = append(e, i)
return nil
})
require.Nil(t, err)
require.Equal(t, 10, len(e))
require.Equal(t, []int{2, 3, 4, 5, 6, 7, 8, 9, 10, 11}, e)
// should be able to get length
require.Equal(t, 10, c.Len())
// should be able to get values
require.Equal(t, 2, c.Get(0))
require.Equal(t, 3, c.Get(1))
// slice should return all values
require.Equal(t, []int{2, 3, 4, 5, 6, 7, 8, 9, 10, 11}, c.Slice())
c.Add(12)
c.Add(13)
// should be able to range over values
e = []int{}
err = c.Range(func(i int) error {
e = append(e, i)
return nil
})
require.Nil(t, err)
require.Equal(t, 10, len(e))
require.Equal(t, []int{4, 5, 6, 7, 8, 9, 10, 11, 12, 13}, e)
require.Equal(t, 4, c.Get(0))
require.Equal(t, 5, c.Get(1))
// slice should return all values
require.Equal(t, []int{4, 5, 6, 7, 8, 9, 10, 11, 12, 13}, c.Slice())
}
@@ -29,26 +29,21 @@ import (
var _ entity.EntityStoreServer = &sqlEntityServer{}
func ProvideSQLEntityServer(db db.EntityDBInterface /*, cfg *setting.Cfg */) (entity.EntityStoreServer, error) {
snode, err := snowflake.NewNode(rand.Int63n(1024))
if err != nil {
return nil, err
}
entityServer := &sqlEntityServer{
db: db,
log: log.New("sql-entity-server"),
snowflake: snode,
db: db,
log: log.New("sql-entity-server"),
}
return entityServer, nil
}
type sqlEntityServer struct {
log log.Logger
db db.EntityDBInterface // needed to keep xorm engine in scope
sess *session.SessionDB
dialect migrator.Dialect
snowflake *snowflake.Node
log log.Logger
db db.EntityDBInterface // needed to keep xorm engine in scope
sess *session.SessionDB
dialect migrator.Dialect
snowflake *snowflake.Node
broadcaster Broadcaster[*entity.Entity]
}
func (s *sqlEntityServer) Init() error {
@@ -77,6 +72,24 @@ func (s *sqlEntityServer) Init() error {
s.sess = sess
s.dialect = migrator.NewDialect(engine.DriverName())
// initialize snowflake generator
s.snowflake, err = snowflake.NewNode(rand.Int63n(1024))
if err != nil {
return err
}
// set up the broadcaster
s.broadcaster, err = NewBroadcaster(context.Background(), func(stream chan *entity.Entity) error {
// start the poller
go s.poller(stream)
return nil
})
if err != nil {
return err
}
return nil
}
@@ -92,6 +105,7 @@ func (s *sqlEntityServer) getReadFields(r *entity.ReadEntityRequest) []string {
"meta",
"title", "slug", "description", "labels", "fields",
"message",
"action",
}
if r.WithBody {
@@ -138,6 +152,7 @@ func (s *sqlEntityServer) rowToEntity(ctx context.Context, rows *sql.Rows, r *en
&raw.Meta,
&raw.Title, &raw.Slug, &raw.Description, &labels, &fields,
&raw.Message,
&raw.Action,
}
if r.WithBody {
args = append(args, &raw.Body)
@@ -423,7 +438,7 @@ func (s *sqlEntityServer) Create(ctx context.Context, r *entity.CreateEntityRequ
current.Message = r.Entity.Message
}
// Update version
// Update resource version
current.ResourceVersion = s.snowflake.Generate().Int64()
values := map[string]any{
@@ -455,6 +470,7 @@ func (s *sqlEntityServer) Create(ctx context.Context, r *entity.CreateEntityRequ
"origin_key": current.Origin.Key,
"origin_ts": current.Origin.Time,
"message": current.Message,
"action": entity.Entity_CREATED,
}
// 1. Add row to the `entity_history` values
@@ -627,7 +643,7 @@ func (s *sqlEntityServer) Update(ctx context.Context, r *entity.UpdateEntityRequ
current.Message = r.Entity.Message
}
// Update version
// Update resource version
current.ResourceVersion = s.snowflake.Generate().Int64()
values := map[string]any{
@@ -661,6 +677,7 @@ func (s *sqlEntityServer) Update(ctx context.Context, r *entity.UpdateEntityRequ
"origin_key": current.Origin.Key,
"origin_ts": current.Origin.Time,
"message": current.Message,
"action": entity.Entity_UPDATED,
}
// 1. Add the `entity_history` values
@@ -680,6 +697,7 @@ func (s *sqlEntityServer) Update(ctx context.Context, r *entity.UpdateEntityRequ
delete(values, "name")
delete(values, "created_at")
delete(values, "created_by")
delete(values, "action")
err = s.dialect.Update(
ctx,
@@ -792,13 +810,68 @@ func (s *sqlEntityServer) Delete(ctx context.Context, r *entity.DeleteEntityRequ
}
func (s *sqlEntityServer) doDelete(ctx context.Context, tx *session.SessionTx, ent *entity.Entity) error {
_, err := tx.Exec(ctx, "DELETE FROM entity WHERE guid=?", ent.Guid)
// Update resource version
ent.ResourceVersion = s.snowflake.Generate().Int64()
labels, err := json.Marshal(ent.Labels)
if err != nil {
s.log.Error("error marshalling labels", "msg", err.Error())
return err
}
// TODO: keep history? would need current version bump, and the "write" would have to get from history
_, err = tx.Exec(ctx, "DELETE FROM entity_history WHERE guid=?", ent.Guid)
fields, err := json.Marshal(ent.Fields)
if err != nil {
s.log.Error("error marshalling fields", "msg", err.Error())
return err
}
errors, err := json.Marshal(ent.Errors)
if err != nil {
s.log.Error("error marshalling errors", "msg", err.Error())
return err
}
values := map[string]any{
// below are only set in history table
"guid": ent.Guid,
"key": ent.Key,
"namespace": ent.Namespace,
"group": ent.Group,
"resource": ent.Resource,
"name": ent.Name,
"created_at": ent.CreatedAt,
"created_by": ent.CreatedBy,
// below are updated
"group_version": ent.GroupVersion,
"folder": ent.Folder,
"slug": ent.Slug,
"updated_at": ent.UpdatedAt,
"updated_by": ent.UpdatedBy,
"body": ent.Body,
"meta": ent.Meta,
"status": ent.Status,
"size": ent.Size,
"etag": ent.ETag,
"resource_version": ent.ResourceVersion,
"title": ent.Title,
"description": ent.Description,
"labels": labels,
"fields": fields,
"errors": errors,
"origin": ent.Origin.Source,
"origin_key": ent.Origin.Key,
"origin_ts": ent.Origin.Time,
"message": ent.Message,
"action": entity.Entity_DELETED,
}
// 1. Add the `entity_history` values
if err := s.dialect.Insert(ctx, tx, "entity_history", values); err != nil {
s.log.Error("error inserting entity history", "msg", err.Error())
return err
}
_, err = tx.Exec(ctx, "DELETE FROM entity WHERE guid=?", ent.Guid)
if err != nil {
return err
}
@@ -1106,12 +1179,356 @@ func (s *sqlEntityServer) List(ctx context.Context, r *entity.EntityListRequest)
return rsp, err
}
func (s *sqlEntityServer) Watch(*entity.EntityWatchRequest, entity.EntityStore_WatchServer) error {
func (s *sqlEntityServer) Watch(r *entity.EntityWatchRequest, w entity.EntityStore_WatchServer) error {
if err := s.Init(); err != nil {
return err
}
return fmt.Errorf("unimplemented")
user, err := appcontext.User(w.Context())
if err != nil {
return err
}
if user == nil {
return fmt.Errorf("missing user in context")
}
// collect and send any historical events
err = s.watchInit(w.Context(), r, w)
if err != nil {
return err
}
// subscribe to new events
err = s.watch(w.Context(), r, w)
if err != nil {
s.log.Error("watch error", "err", err)
return err
}
return nil
}
// watchInit is a helper function to send the initial set of entities to the client
func (s *sqlEntityServer) watchInit(ctx context.Context, r *entity.EntityWatchRequest, w entity.EntityStore_WatchServer) error {
rr := &entity.ReadEntityRequest{
WithBody: r.WithBody,
WithStatus: r.WithStatus,
}
fields := s.getReadFields(rr)
entityQuery := selectQuery{
dialect: s.dialect,
fields: fields,
from: "entity", // the table
args: []any{},
limit: 100, // r.Limit,
oneExtra: true, // request one more than the limit (and show next token if it exists)
}
// if we got an initial resource version, start from that location in the history
fromZero := true
if r.Since > 0 {
entityQuery.from = "entity_history"
entityQuery.addWhere("resource_version > ?", r.Since)
fromZero = false
}
// TODO fix this
// entityQuery.addWhere("namespace", user.OrgID)
if len(r.Resource) > 0 {
entityQuery.addWhereIn("resource", r.Resource)
}
if len(r.Key) > 0 {
where := []string{}
args := []any{}
for _, k := range r.Key {
key, err := entity.ParseKey(k)
if err != nil {
return err
}
args = append(args, key.Namespace, key.Group, key.Resource)
whereclause := "(" + s.dialect.Quote("namespace") + "=? AND " + s.dialect.Quote("group") + "=? AND " + s.dialect.Quote("resource") + "=?"
if key.Name != "" {
args = append(args, key.Name)
whereclause += " AND " + s.dialect.Quote("name") + "=?"
}
whereclause += ")"
where = append(where, whereclause)
}
entityQuery.addWhere("("+strings.Join(where, " OR ")+")", args...)
}
// Folder guid
if r.Folder != "" {
entityQuery.addWhere("folder", r.Folder)
}
if len(r.Labels) > 0 {
var args []any
var conditions []string
for labelKey, labelValue := range r.Labels {
args = append(args, labelKey)
args = append(args, labelValue)
conditions = append(conditions, "(label = ? AND value = ?)")
}
query := "SELECT guid FROM entity_labels" +
" WHERE (" + strings.Join(conditions, " OR ") + ")" +
" GROUP BY guid" +
" HAVING COUNT(label) = ?"
args = append(args, len(r.Labels))
entityQuery.addWhereInSubquery("guid", query, args)
}
entityQuery.addOrderBy("resource_version", Ascending)
var err error
s.log.Debug("watch init", "since", r.Since)
for hasmore := true; hasmore; {
err = func() error {
query, args := entityQuery.toQuery()
rows, err := s.sess.Query(ctx, query, args...)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
found := int64(0)
for rows.Next() {
found++
if found > entityQuery.limit {
entityQuery.offset += entityQuery.limit
return nil
}
result, err := s.rowToEntity(ctx, rows, rr)
if err != nil {
return err
}
if result.ResourceVersion > r.Since {
r.Since = result.ResourceVersion
}
if fromZero {
result.Action = entity.Entity_CREATED
}
s.log.Debug("sending init event", "guid", result.Guid, "action", result.Action, "rv", result.ResourceVersion)
err = w.Send(&entity.EntityWatchResponse{
Timestamp: time.Now().UnixMilli(),
Entity: result,
})
if err != nil {
return err
}
}
hasmore = false
return nil
}()
if err != nil {
return err
}
}
return nil
}
func (s *sqlEntityServer) poller(stream chan *entity.Entity) {
var err error
since := s.snowflake.Generate().Int64()
t := time.NewTicker(5 * time.Second)
defer t.Stop()
for range t.C {
since, err = s.poll(context.Background(), since, stream)
if err != nil {
s.log.Error("watch error", "err", err)
}
}
}
func (s *sqlEntityServer) poll(ctx context.Context, since int64, out chan *entity.Entity) (int64, error) {
s.log.Debug("watch poll", "since", since)
rr := &entity.ReadEntityRequest{
WithBody: true,
WithStatus: true,
}
fields := s.getReadFields(rr)
for hasmore := true; hasmore; {
err := func() error {
entityQuery := selectQuery{
dialect: s.dialect,
fields: fields,
from: "entity_history", // the table
args: []any{},
limit: 100, // r.Limit,
// offset: 0,
oneExtra: true, // request one more than the limit (and show next token if it exists)
orderBy: []string{"resource_version"},
}
entityQuery.addWhere("resource_version > ?", since)
query, args := entityQuery.toQuery()
rows, err := s.sess.Query(ctx, query, args...)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
found := int64(0)
for rows.Next() {
found++
if found > entityQuery.limit {
return nil
}
result, err := s.rowToEntity(ctx, rows, rr)
if err != nil {
return err
}
if result.ResourceVersion > since {
since = result.ResourceVersion
}
s.log.Debug("sending poll result", "guid", result.Guid, "action", result.Action, "rv", result.ResourceVersion)
out <- result
}
hasmore = false
return nil
}()
if err != nil {
return since, err
}
}
return since, nil
}
func watchMatches(r *entity.EntityWatchRequest, result *entity.Entity) bool {
// Resource version too old
if result.ResourceVersion <= r.Since {
return false
}
// Folder guid
if r.Folder != "" && r.Folder != result.Folder {
return false
}
// must match at least one resource if specified
if len(r.Resource) > 0 {
matched := false
for _, res := range r.Resource {
if res == result.Resource {
matched = true
break
}
}
if !matched {
return false
}
}
// must match at least one key if specified
if len(r.Key) > 0 {
matched := false
for _, k := range r.Key {
key, err := entity.ParseKey(k)
if err != nil {
return false
}
if key.Namespace == result.Namespace && key.Group == result.Group && key.Resource == result.Resource && (key.Name == "" || key.Name == result.Name) {
matched = true
break
}
}
if !matched {
return false
}
}
// must match at least one label/value pair if specified
// TODO should this require matching all label conditions?
if len(r.Labels) > 0 {
matched := false
for labelKey, labelValue := range r.Labels {
if result.Labels[labelKey] == labelValue {
matched = true
break
}
}
if !matched {
return false
}
}
return true
}
// watch is a helper to get the next set of entities and send them to the client
func (s *sqlEntityServer) watch(ctx context.Context, r *entity.EntityWatchRequest, w entity.EntityStore_WatchServer) error {
s.log.Debug("watch started", "since", r.Since)
evts, err := s.broadcaster.Subscribe(w.Context())
if err != nil {
return err
}
for {
select {
// user closed the connection
case <-w.Context().Done():
return nil
// got a raw result from the broadcaster
case result := <-evts:
// result doesn't match our watch params, skip it
if !watchMatches(r, result) {
s.log.Debug("watch result not matched", "guid", result.Guid, "action", result.Action, "rv", result.ResourceVersion)
break
}
// remove the body and status if not requested
if !r.WithBody {
result.Body = nil
}
if !r.WithStatus {
result.Status = nil
}
// update r.Since value so we don't send earlier results again
r.Since = result.ResourceVersion
s.log.Debug("sending watch result", "guid", result.Guid, "action", result.Action, "rv", result.ResourceVersion)
err = w.Send(&entity.EntityWatchResponse{
Timestamp: time.Now().UnixMilli(),
Entity: result,
})
if err != nil {
return err
}
}
}
}
func (s *sqlEntityServer) FindReferences(ctx context.Context, r *entity.ReferenceRequest) (*entity.EntityListResponse, error) {