285 lines
10 KiB
Go
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
|
|
}
|