Files
rvbox/internal/server/store/query.go
T

285 lines
10 KiB
Go

package store
import (
"context"
"database/sql"
"errors"
"time"
"github.com/rvbox/rvbox/internal/domain"
)
// ClientView is the read-only, protocol-neutral representation of a client.
// The store deliberately returns copies of all byte slices so callers cannot
// mutate memory owned by a database driver or a shared scan buffer.
type ClientView struct {
ClientID string
Connected bool
ConnectedAt *time.Time
LastSeenAt *time.Time
RunningCommands uint32
QueuedCommands uint32
Platform uint32
Architecture string
DaemonVersion string
DaemonCWD string
SupportedShells []byte
ClientInstanceID [16]byte
PendingInstanceID *[16]byte
PendingInstanceSeenAt *time.Time
}
// CommandView is the durable command metadata exposed to the control layer.
// ExecutionSpec is returned in decoded form; the compressed representation
// never crosses the store boundary.
type CommandView struct {
IssueUUID domain.UUID
ClientID string
IssueTime time.Time
ServerReceiptTime time.Time
QueueExpiryTime *time.Time
TerminalTime *time.Time
Lifecycle uint32
Revision uint64
LastEventSeq uint64
ExitCode *int32
OutputTruncated bool
OutputIncomplete bool
RetainedCompressedBytes uint64
ExecutionSpec []byte
WindowsIdentity []byte
}
// ListClientViews returns clients ordered by client_id. The afterClientID
// value is an exclusive lexical cursor; an empty value starts at the first
// row. The boolean reports whether another row exists.
func (store *Store) ListClientViews(ctx context.Context, afterClientID string, limit uint32) ([]ClientView, bool, error) {
if limit == 0 || limit > 1000 {
return nil, false, errors.New("invalid client page size")
}
database, err := store.openDatabase()
if err != nil {
return nil, false, err
}
rows, err := database.QueryContext(ctx, `SELECT
c.client_id, c.platform, c.architecture, c.daemon_version, c.daemon_cwd,
c.supported_shells, c.client_instance_id, c.connected_at, c.last_seen_at,
c.pending_instance_id, c.pending_instance_seen_at,
EXISTS(SELECT 1 FROM sessions s WHERE s.client_id = c.client_id AND s.closed_at IS NULL AND s.fenced_at IS NULL),
(SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (3,4)),
(SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (1,2))
FROM clients c WHERE c.client_id > ? ORDER BY c.client_id LIMIT ?`, afterClientID, limit+1)
if err != nil {
return nil, false, err
}
defer rows.Close()
clients := make([]ClientView, 0, limit)
for rows.Next() {
view, err := scanClientView(rows)
if err != nil {
return nil, false, err
}
if uint32(len(clients)) < limit {
clients = append(clients, view)
} else {
return clients, true, rows.Err()
}
}
if err := rows.Err(); err != nil {
return nil, false, err
}
return clients, false, nil
}
// GetClientView returns one client or ErrClientNotFound. Connection state is
// derived from the live-session row rather than the historical timestamps.
func (store *Store) GetClientView(ctx context.Context, clientID string) (ClientView, error) {
if clientID == "" {
return ClientView{}, ErrClientNotFound
}
database, err := store.openDatabase()
if err != nil {
return ClientView{}, err
}
row := database.QueryRowContext(ctx, `SELECT
c.client_id, c.platform, c.architecture, c.daemon_version, c.daemon_cwd,
c.supported_shells, c.client_instance_id, c.connected_at, c.last_seen_at,
c.pending_instance_id, c.pending_instance_seen_at,
EXISTS(SELECT 1 FROM sessions s WHERE s.client_id = c.client_id AND s.closed_at IS NULL AND s.fenced_at IS NULL),
(SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (3,4)),
(SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (1,2))
FROM clients c WHERE c.client_id = ?`, clientID)
view, err := scanClientView(row)
if errors.Is(err, sql.ErrNoRows) {
return ClientView{}, ErrClientNotFound
}
return view, err
}
func scanClientView(scanner interface{ Scan(...any) error }) (ClientView, error) {
var view ClientView
var instance, pending []byte
var connectedAt, lastSeen, pendingSeen sql.NullInt64
var connected, running, queued int64
if err := scanner.Scan(&view.ClientID, &view.Platform, &view.Architecture, &view.DaemonVersion, &view.DaemonCWD,
&view.SupportedShells, &instance, &connectedAt, &lastSeen, &pending, &pendingSeen,
&connected, &running, &queued); err != nil {
return view, err
}
if len(instance) != 16 || (pending != nil && len(pending) != 16) {
return view, ErrInvalidSegmentRecord
}
copy(view.ClientInstanceID[:], instance)
if pending != nil {
var value [16]byte
copy(value[:], pending)
view.PendingInstanceID = &value
}
view.Connected = connected == 1
if running < 0 || running > int64(^uint32(0)) || queued < 0 || queued > int64(^uint32(0)) {
return view, ErrInvalidSegmentRecord
}
view.RunningCommands, view.QueuedCommands = uint32(running), uint32(queued)
view.ConnectedAt = nullableTime(connectedAt)
view.LastSeenAt = nullableTime(lastSeen)
view.PendingInstanceSeenAt = nullableTime(pendingSeen)
view.SupportedShells = append([]byte(nil), view.SupportedShells...)
return view, nil
}
func nullableTime(value sql.NullInt64) *time.Time {
if !value.Valid {
return nil
}
instant := time.Unix(0, value.Int64).UTC()
return &instant
}
// CommandPage describes a stable descending command cursor. SnapshotBoundary
// is an inclusive issue-time ceiling captured by the first page. AfterTime and
// AfterUUID are the exclusive position from the previous page.
type CommandPage struct {
ClientID string
IncludeTerminal bool
Limit uint32
SnapshotBoundary int64
AfterTime int64
AfterUUID domain.UUID
HasAfter bool
}
// ListCommandViews reads a stable, descending command page. The caller must
// preserve SnapshotBoundary and the returned final (issue_time, UUID) pair in
// its authenticated cursor.
func (store *Store) ListCommandViews(ctx context.Context, page CommandPage) ([]CommandView, bool, error) {
if page.ClientID == "" || page.Limit == 0 || page.Limit > 1000 || page.SnapshotBoundary <= 0 {
return nil, false, errors.New("invalid command page")
}
database, err := store.openDatabase()
if err != nil {
return nil, false, err
}
query := `SELECT issue_uuid, client_id, issue_time, server_receipt_time,
queue_expiry_time, terminal_time, lifecycle, revision, last_event_seq,
exit_code, output_truncated, output_incomplete, retained_compressed_bytes,
execution_spec, execution_spec_raw_bytes, windows_execution_identity
FROM commands WHERE client_id = ? AND issue_time <= ?`
args := []any{page.ClientID, page.SnapshotBoundary}
if !page.IncludeTerminal {
query += ` AND lifecycle NOT BETWEEN 5 AND 11`
}
if page.HasAfter {
query += ` AND (issue_time < ? OR (issue_time = ? AND issue_uuid < ?))`
args = append(args, page.AfterTime, page.AfterTime, page.AfterUUID[:])
}
query += ` ORDER BY issue_time DESC, issue_uuid DESC LIMIT ?`
args = append(args, page.Limit+1)
rows, err := database.QueryContext(ctx, query, args...)
if err != nil {
return nil, false, err
}
defer rows.Close()
commands := make([]CommandView, 0, page.Limit)
for rows.Next() {
view, err := scanCommandView(rows)
if err != nil {
return nil, false, err
}
if uint32(len(commands)) < page.Limit {
commands = append(commands, view)
} else {
return commands, true, rows.Err()
}
}
if err := rows.Err(); err != nil {
return nil, false, err
}
return commands, false, nil
}
// GetCommandView returns one retained command, optionally constrained to its
// client. Evicted commands are represented by ErrCommandNotFound; the
// tombstone remains available to reconciliation rather than control history.
func (store *Store) GetCommandView(ctx context.Context, clientID string, issue domain.UUID) (CommandView, error) {
if issue == (domain.UUID{}) {
return CommandView{}, ErrCommandNotFound
}
database, err := store.openDatabase()
if err != nil {
return CommandView{}, err
}
query := `SELECT issue_uuid, client_id, issue_time, server_receipt_time,
queue_expiry_time, terminal_time, lifecycle, revision, last_event_seq,
exit_code, output_truncated, output_incomplete, retained_compressed_bytes,
execution_spec, execution_spec_raw_bytes, windows_execution_identity
FROM commands WHERE issue_uuid = ?`
args := []any{issue[:]}
if clientID != "" {
query += ` AND client_id = ?`
args = append(args, clientID)
}
view, err := scanCommandView(database.QueryRowContext(ctx, query, args...))
if errors.Is(err, sql.ErrNoRows) {
return CommandView{}, ErrCommandNotFound
}
return view, err
}
func scanCommandView(scanner interface{ Scan(...any) error }) (CommandView, error) {
var view CommandView
var issue, stored []byte
var issueTime, receipt int64
var expiry, terminal sql.NullInt64
var exit sql.NullInt64
var truncated, incomplete int
var rawBytes uint64
if err := scanner.Scan(&issue, &view.ClientID, &issueTime, &receipt, &expiry, &terminal,
&view.Lifecycle, &view.Revision, &view.LastEventSeq, &exit, &truncated, &incomplete,
&view.RetainedCompressedBytes, &stored, &rawBytes, &view.WindowsIdentity); err != nil {
return view, err
}
if len(issue) != 16 {
return view, ErrInvalidSegmentRecord
}
copy(view.IssueUUID[:], issue)
view.IssueTime = time.Unix(0, issueTime).UTC()
view.ServerReceiptTime = time.Unix(0, receipt).UTC()
view.QueueExpiryTime = nullableTime(expiry)
view.TerminalTime = nullableTime(terminal)
view.OutputTruncated, view.OutputIncomplete = truncated == 1, incomplete == 1
if exit.Valid {
if exit.Int64 < -1<<31 || exit.Int64 > 1<<31-1 {
return view, ErrInvalidSegmentRecord
}
value := int32(exit.Int64)
view.ExitCode = &value
}
var err error
view.ExecutionSpec, err = decompressCommandSpec(stored, rawBytes)
if err != nil {
return view, err
}
view.WindowsIdentity = append([]byte(nil), view.WindowsIdentity...)
return view, nil
}