Clean up and comments before merge.

This commit is contained in:
Zachary Wasserman
2020-07-21 14:05:46 -07:00
parent 96fc090723
commit 7494513400
23 changed files with 259 additions and 272 deletions
-1
View File
@@ -502,7 +502,6 @@ func testHostAdditional(t *testing.T, ds kolide.Datastore) {
DetailUpdateTime: time.Now(),
LabelUpdateTime: time.Now(),
SeenTime: time.Now(),
LabelUpdateTime: time.Now(),
OsqueryHostID: "foobar",
NodeKey: "nodekey",
UUID: "uuid",
-1
View File
@@ -70,7 +70,6 @@ var testFunctions = [...]func(*testing.T, kolide.Datastore){
testCountHostsInTargets,
testHostStatus,
testHostIDsInTargets,
testResetOptions,
testApplyOsqueryOptions,
testApplyOsqueryOptionsNoOverrides,
testOsqueryOptionsForHost,
+3 -5
View File
@@ -20,7 +20,7 @@ func (d *Datastore) ApplyLabelSpecs(specs []*kolide.LabelSpec) (err error) {
platform,
label_type,
label_membership_type
) VALUES ( ?, ?, ?, ?, ? , ?)
) VALUES ( ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
name = VALUES(name),
description = VALUES(description),
@@ -54,8 +54,7 @@ func (d *Datastore) ApplyLabelSpecs(specs []*kolide.LabelSpec) (err error) {
sql = `
SELECT id from labels WHERE name = ?
`
err = tx.Get(&labelID, sql, s.Name)
if err != nil {
if err := tx.Get(&labelID, sql, s.Name); err != nil {
return errors.Wrap(err, "get label ID")
}
@@ -121,8 +120,7 @@ func (d *Datastore) GetLabelSpecs() ([]*kolide.LabelSpec, error) {
for _, spec := range specs {
if spec.LabelType != kolide.LabelTypeBuiltIn &&
spec.LabelMembershipType == kolide.LabelMembershipTypeManual {
err := d.getLabelHostnames(spec)
if err != nil {
if err := d.getLabelHostnames(spec); err != nil {
return nil, err
}
}
+30
View File
@@ -0,0 +1,30 @@
package mysql
import (
"strconv"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestBatchHostnamesSmall(t *testing.T) {
small := []string{"foo", "bar", "baz"}
batched := batchHostnames(small)
require.Equal(t, 1, len(batched))
assert.Equal(t, small, batched[0])
}
func TestBatchHostnamesLarge(t *testing.T) {
large := []string{}
for i := 0; i < 230000; i++ {
large = append(large, strconv.Itoa(i))
}
batched := batchHostnames(large)
require.Equal(t, 5, len(batched))
assert.Equal(t, large[:50000], batched[0])
assert.Equal(t, large[50000:100000], batched[1])
assert.Equal(t, large[100000:150000], batched[2])
assert.Equal(t, large[150000:200000], batched[3])
assert.Equal(t, large[200000:230000], batched[4])
}
+73 -28
View File
@@ -3,41 +3,86 @@ package live_query
import (
"testing"
"github.com/kolide/fleet/server/kolide"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMapBitfield(t *testing.T) {
// empty
assert.Equal(t, []byte{}, mapBitfield(nil))
assert.Equal(t, []byte{}, mapBitfield([]uint{}))
var testFunctions = [...]func(*testing.T, kolide.LiveQueryStore){
testLiveQuery,
testLiveQueryNoTargets,
testLiveQueryStopQuery,
}
// one byte
assert.Equal(t, []byte("\x80"), mapBitfield([]uint{0}))
assert.Equal(t, []byte("\x40"), mapBitfield([]uint{1}))
assert.Equal(t, []byte("\xc0"), mapBitfield([]uint{0, 1}))
func testLiveQuery(t *testing.T, store kolide.LiveQueryStore) {
queries, err := store.QueriesForHost(1)
assert.NoError(t, err)
assert.Len(t, queries, 0)
queries, err = store.QueriesForHost(3)
assert.NoError(t, err)
assert.Len(t, queries, 0)
assert.Equal(t, []byte("\x08"), mapBitfield([]uint{4}))
assert.Equal(t, []byte("\xf8"), mapBitfield([]uint{0, 1, 2, 3, 4}))
assert.Equal(t, []byte("\xff"), mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7}))
assert.NoError(t, store.RunQuery("test", "select 1", []uint{1, 3}))
assert.NoError(t, store.RunQuery("test2", "select 2", []uint{3}))
assert.NoError(t, store.RunQuery("test3", "select 3", []uint{1}))
assert.NoError(t, store.RunQuery("test4", "select 4", []uint{4}))
// two bytes
assert.Equal(t, []byte("\x00\x80"), mapBitfield([]uint{8}))
assert.Equal(t, []byte("\xff\x80"), mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7, 8}))
// more bytes
assert.Equal(
t,
[]byte("\xff\x80\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00 "),
mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7, 8, 170}),
queries, err = store.QueriesForHost(1)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test": "select 1",
"test3": "select 3",
},
queries,
)
assert.Equal(
t,
[]byte("\xff\x80\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00@\x00\x00\x00\x00\x00\x00 "),
mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7, 8, 113, 170}),
queries, err = store.QueriesForHost(2)
assert.NoError(t, err)
assert.Len(t, queries, 0)
queries, err = store.QueriesForHost(3)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test": "select 1",
"test2": "select 2",
},
queries,
)
assert.Equal(
t,
[]byte("\x00\x00\x00\x00\x00\x00\x00\x00\x00\x01"),
mapBitfield([]uint{79}),
assert.NoError(t, store.QueryCompletedByHost("test", 1))
assert.NoError(t, store.QueryCompletedByHost("test2", 3))
queries, err = store.QueriesForHost(1)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test3": "select 3",
},
queries,
)
queries, err = store.QueriesForHost(2)
assert.NoError(t, err)
assert.Len(t, queries, 0)
queries, err = store.QueriesForHost(3)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test": "select 1",
},
queries,
)
}
func testLiveQueryNoTargets(t *testing.T, store kolide.LiveQueryStore) {
assert.Error(t, store.RunQuery("test", "select 1", []uint{}))
}
func testLiveQueryStopQuery(t *testing.T, store kolide.LiveQueryStore) {
require.NoError(t, store.RunQuery("test", "select 1", []uint{1, 3}))
require.NoError(t, store.RunQuery("test2", "select 2", []uint{1, 3}))
require.NoError(t, store.StopQuery("test"))
queries, err := store.QueriesForHost(1)
require.NoError(t, err)
assert.Len(t, queries, 1)
}
+76 -47
View File
@@ -1,7 +1,29 @@
// package live_query implements an interface for storing and
// retrieving live queries.
//
// Design
//
// This package operates by storing a single redis key for host
// targeting information. This key has a known prefix, and the data
// is a bitfield representing _all_ the hosts in fleet.
//
// In this model, a live query creation is a few redis writes. While a
// host checkin needs to scan the keyspace for matching key, and then
// fetch the bitfield value for their id. While this scan might be
// expensive, this model fits very well with having a lot of hosts and
// very few live queries.
//
// A contrasting model, for the case of fewer hosts, but a lot of live
// queries, is to have a set per host. In this case, the LQ is pushed
// into each host's set. This model has many potential writes for LQ
// creation, but a host checkin has very few.
//
// We believe that normal fleet usage has many hosts, and a small
// number of live queries targeting all of them. This was a big
// factor in choosing this implementation.
package live_query
import (
"fmt"
"strings"
"time"
@@ -11,7 +33,8 @@ import (
const (
bitsInByte = 8
queryKeyPrefix = "query:"
queryKeyPrefix = "livequery:"
sqlKeyPrefix = "sql:"
queryExpiration = 7 * 24 * time.Hour
)
@@ -26,6 +49,10 @@ func NewRedisLiveQuery(pool *redis.Pool) *redisLiveQuery {
return &redisLiveQuery{pool: pool}
}
func generateKeys(name string) (targetsKey, sqlKey string) {
return queryKeyPrefix + name, sqlKeyPrefix + queryKeyPrefix + name
}
func (r *redisLiveQuery) RunQuery(name, sql string, hostIDs []uint) error {
if len(hostIDs) == 0 {
return errors.New("no hosts targeted")
@@ -34,14 +61,22 @@ func (r *redisLiveQuery) RunQuery(name, sql string, hostIDs []uint) error {
conn := r.pool.Get()
defer conn.Close()
// Map the targeted host IDs to a bitfield and store in a key containing the
// query anme and SQL.
key := fmt.Sprintf(queryKeyPrefix+"%s:%s", name, sql)
bitfield := mapBitfield(hostIDs)
_, err := conn.Do("SET", key, bitfield, "EX", queryExpiration.Seconds())
// Map the targeted host IDs to a bitfield. Store targets in one key and SQL
// in another.
targetKey, sqlKey := generateKeys(name)
targets := mapBitfield(hostIDs)
// Ensure to set SQL first or else we can end up in a weird state in which a
// client reads that the query exists but cannot look up the SQL.
err := conn.Send("SET", sqlKey, sql, "EX", queryExpiration.Seconds())
if err != nil {
return errors.Wrap(err, "set query in Redis")
return errors.Wrap(err, "set sql")
}
_, err = conn.Do("SET", targetKey, targets, "EX", queryExpiration.Seconds())
if err != nil {
return errors.Wrap(err, "set targets")
}
return nil
}
@@ -49,22 +84,9 @@ func (r *redisLiveQuery) StopQuery(name string) error {
conn := r.pool.Get()
defer conn.Close()
// Find key for this query.
keys, err := scanKeys(conn, queryKeyPrefix+name+":*")
if err != nil {
return errors.Wrap(err, "scan for query key")
}
if len(keys) == 0 {
return errors.Errorf("query %s not found", name)
}
if len(keys) > 1 {
return errors.Errorf("found more than one query matching %s", name)
}
// Set the bitfield for this host.
key := keys[0]
if _, err := conn.Do("DEL", key); err != nil {
return errors.Wrap(err, "del query key")
targetKey, sqlKey := generateKeys(name)
if _, err := conn.Do("DEL", targetKey, sqlKey); err != nil {
return errors.Wrap(err, "del query keys")
}
return nil
@@ -84,7 +106,15 @@ func (r *redisLiveQuery) QueriesForHost(hostID uint) (map[string]string, error)
// targets of the query.
for _, key := range queryKeys {
if err := conn.Send("GETBIT", key, hostID); err != nil {
return nil, errors.Wrap(err, "getbit query key")
return nil, errors.Wrap(err, "getbit query targets")
}
// Additionally get SQL even though we don't yet know whether this query
// is targeted to the host. This allows us to avoid an additional
// roundtrip to the Redis server and likely has little cost due to the
// small number of queries and limited size of SQL
if err = conn.Send("GET", sqlKeyPrefix+key); err != nil {
return nil, errors.Wrap(err, "get query sql")
}
}
@@ -93,24 +123,34 @@ func (r *redisLiveQuery) QueriesForHost(hostID uint) (map[string]string, error)
return nil, errors.Wrap(err, "flush pipeline")
}
// Receive target information in order of pipelined calls.
// Receive target and SQL in order of pipelined calls.
queries := make(map[string]string)
for _, key := range queryKeys {
name := strings.TrimPrefix(key, queryKeyPrefix)
targeted, err := redis.Int(conn.Receive())
if err != nil {
return nil, errors.Wrap(err, "receive int")
return nil, errors.Wrap(err, "receive target")
}
// Be sure to read SQL even if we are not going to include this query.
// Otherwise we will read an incorrect number of returned results from
// the pipeline.
sql, err := redis.String(conn.Receive())
if err != nil {
// Not being able to get the sql for a matched could mean things
// have ended up in a weird state. Or it could be that the query was
// stopped since we did the key scan. In any case, attempt to clean
// up here.
_ = r.StopQuery(name)
return nil, errors.Wrap(err, "receive sql")
}
if targeted == 0 {
// Host not targeted with this query
continue
}
// Split the key to get the query name and SQL
splits := strings.SplitN(key, ":", 3)
if len(splits) != 3 {
return nil, errors.Errorf("query key did not have 3 components: %s", key)
}
name, sql := splits[1], splits[2]
queries[name] = sql
}
@@ -121,21 +161,10 @@ func (r *redisLiveQuery) QueryCompletedByHost(name string, hostID uint) error {
conn := r.pool.Get()
defer conn.Close()
// Find key for this query.
keys, err := scanKeys(conn, queryKeyPrefix+name+":*")
if err != nil {
return errors.Wrap(err, "scan for query key")
}
if len(keys) == 0 {
return errors.Errorf("query %s not found", name)
}
if len(keys) > 1 {
return errors.Errorf("found more than one query matching %s", name)
}
targetKey, _ := generateKeys(name)
// Set the bitfield for this host.
key := keys[0]
if _, err := conn.Do("SETBIT", key, hostID, 0); err != nil {
// Update the bitfield for this host.
if _, err := conn.Do("SETBIT", targetKey, hostID, 0); err != nil {
return errors.Wrap(err, "setbit query key")
}
+29 -73
View File
@@ -5,7 +5,6 @@ import (
"os"
"testing"
"github.com/kolide/fleet/server/kolide"
"github.com/kolide/fleet/server/pubsub"
"github.com/kolide/fleet/server/test"
"github.com/stretchr/testify/assert"
@@ -14,7 +13,7 @@ import (
func TestRedisLiveQuery(t *testing.T) {
if _, ok := os.LookupEnv("REDIS_TEST"); !ok {
t.SkipNow()
t.Skip("Redis tests not requested. Skipping.")
}
for _, f := range testFunctions {
@@ -26,12 +25,6 @@ func TestRedisLiveQuery(t *testing.T) {
}
}
var testFunctions = [...]func(*testing.T, kolide.LiveQueryStore){
testRedisLiveQuery,
testRedisLiveQueryNoTargets,
testRedisLiveQueryStopQuery,
}
func setupRedisLiveQuery(t *testing.T) (store *redisLiveQuery, teardown func()) {
var (
addr = "127.0.0.1:6379"
@@ -45,7 +38,7 @@ func setupRedisLiveQuery(t *testing.T) (store *redisLiveQuery, teardown func())
store = NewRedisLiveQuery(pubsub.NewRedisPool(addr, password))
_, err := store.pool.Get().Do("PING")
require.Nil(t, err)
require.NoError(t, err)
teardown = func() {
store.pool.Get().Do("FLUSHDB")
@@ -55,75 +48,38 @@ func setupRedisLiveQuery(t *testing.T) (store *redisLiveQuery, teardown func())
return store, teardown
}
func testRedisLiveQueryNoTargets(t *testing.T, store kolide.LiveQueryStore) {
assert.Error(t, store.RunQuery("test", "select 1", []uint{}))
}
func TestMapBitfield(t *testing.T) {
// empty
assert.Equal(t, []byte{}, mapBitfield(nil))
assert.Equal(t, []byte{}, mapBitfield([]uint{}))
func testRedisLiveQueryStopQuery(t *testing.T, store kolide.LiveQueryStore) {
require.NoError(t, store.RunQuery("test", "select 1", []uint{1, 3}))
require.NoError(t, store.RunQuery("test2", "select 2", []uint{1, 3}))
require.NoError(t, store.StopQuery("test"))
// one byte
assert.Equal(t, []byte("\x80"), mapBitfield([]uint{0}))
assert.Equal(t, []byte("\x40"), mapBitfield([]uint{1}))
assert.Equal(t, []byte("\xc0"), mapBitfield([]uint{0, 1}))
queries, err := store.QueriesForHost(1)
assert.NoError(t, err)
assert.Len(t, queries, 1)
}
assert.Equal(t, []byte("\x08"), mapBitfield([]uint{4}))
assert.Equal(t, []byte("\xf8"), mapBitfield([]uint{0, 1, 2, 3, 4}))
assert.Equal(t, []byte("\xff"), mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7}))
func testRedisLiveQuery(t *testing.T, store kolide.LiveQueryStore) {
queries, err := store.QueriesForHost(1)
assert.NoError(t, err)
assert.Len(t, queries, 0)
queries, err = store.QueriesForHost(3)
assert.NoError(t, err)
assert.Len(t, queries, 0)
// two bytes
assert.Equal(t, []byte("\x00\x80"), mapBitfield([]uint{8}))
assert.Equal(t, []byte("\xff\x80"), mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7, 8}))
assert.NoError(t, store.RunQuery("test", "select 1", []uint{1, 3}))
assert.NoError(t, store.RunQuery("test2", "select 2", []uint{3}))
assert.NoError(t, store.RunQuery("test3", "select 3", []uint{1}))
assert.NoError(t, store.RunQuery("test4", "select 4", []uint{4}))
queries, err = store.QueriesForHost(1)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test": "select 1",
"test3": "select 3",
},
queries,
// more bytes
assert.Equal(
t,
[]byte("\xff\x80\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00 "),
mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7, 8, 170}),
)
queries, err = store.QueriesForHost(2)
assert.NoError(t, err)
assert.Len(t, queries, 0)
queries, err = store.QueriesForHost(3)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test": "select 1",
"test2": "select 2",
},
queries,
assert.Equal(
t,
[]byte("\xff\x80\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00@\x00\x00\x00\x00\x00\x00 "),
mapBitfield([]uint{0, 1, 2, 3, 4, 5, 6, 7, 8, 113, 170}),
)
assert.NoError(t, store.QueryCompletedByHost("test", 1))
assert.NoError(t, store.QueryCompletedByHost("test2", 3))
queries, err = store.QueriesForHost(1)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test3": "select 3",
},
queries,
)
queries, err = store.QueriesForHost(2)
assert.NoError(t, err)
assert.Len(t, queries, 0)
queries, err = store.QueriesForHost(3)
assert.NoError(t, err)
assert.Equal(t,
map[string]string{
"test": "select 1",
},
queries,
assert.Equal(
t,
[]byte("\x00\x00\x00\x00\x00\x00\x00\x00\x00\x01"),
mapBitfield([]uint{79}),
)
}
+18 -13
View File
@@ -6,7 +6,7 @@ import (
"strconv"
"time"
"github.com/igm/sockjs-go/sockjs"
"github.com/igm/sockjs-go/v3/sockjs"
"github.com/kolide/fleet/server/contexts/viewer"
"github.com/kolide/fleet/server/kolide"
"github.com/kolide/fleet/server/websocket"
@@ -127,6 +127,14 @@ func (svc service) StreamCampaignResults(ctx context.Context, conn *websocket.Co
return
}
// Open the channel from which we will receive incoming query results
// (probably from the redis pubsub implementation)
readChan, err := svc.resultStore.ReadChannel(context.Background(), *campaign)
if err != nil {
conn.WriteJSONError(fmt.Sprintf("cannot open read channel for campaign %d ", campaignID))
return
}
// Setting status to running will cause the query to be returned to the
// targets when they check in for their queries
campaign.Status = kolide.QueryRunning
@@ -140,18 +148,10 @@ func (svc service) StreamCampaignResults(ctx context.Context, conn *websocket.Co
// this campaign.
defer func() {
campaign.Status = kolide.QueryComplete
svc.ds.SaveDistributedQueryCampaign(campaign)
svc.liveQueryStore.StopQuery(strconv.Itoa(int(campaign.ID)))
_ = svc.ds.SaveDistributedQueryCampaign(campaign)
_ = svc.liveQueryStore.StopQuery(strconv.Itoa(int(campaign.ID)))
}()
// Open the channel from which we will receive incoming query results
// (probably from the redis pubsub implementation)
readChan, err := svc.resultStore.ReadChannel(context.Background(), *campaign)
if err != nil {
conn.WriteJSONError(fmt.Sprintf("cannot open read channel for campaign %d ", campaignID))
return
}
status := campaignStatus{
Status: campaignStatusPending,
}
@@ -209,7 +209,7 @@ func (svc service) StreamCampaignResults(ctx context.Context, conn *websocket.Co
}
if err := updateStatus(); err != nil {
svc.logger.Log("msg", "error updating status", "err", err)
_ = svc.logger.Log("msg", "error updating status", "err", err)
return
}
@@ -228,8 +228,13 @@ func (svc service) StreamCampaignResults(ctx context.Context, conn *websocket.Co
case kolide.DistributedQueryResult:
mapHostnameRows(res.Host.HostName, res.Rows)
err = conn.WriteJSONMessage("result", res)
if errors.Cause(err) == sockjs.ErrSessionNotOpen {
// return and stop sending the query if the session was closed
// by the client
return
}
if err != nil {
svc.logger.Log("msg", "error writing to channel", "err", err)
_ = svc.logger.Log("msg", "error writing to channel", "err", err)
}
status.ActualResults++
}
+2 -3
View File
@@ -603,9 +603,8 @@ func (svc service) ingestDistributedQuery(host kolide.Host, name string, rows []
return osqueryError{message: "loading orphaned campaign: " + err.Error()}
}
if campaign.Status == kolide.QueryWaiting &&
campaign.CreatedAt.Before(svc.clock.Now().Add(-1*time.Minute)) {
// Give the client one minute to connect before considering the
if campaign.CreatedAt.Before(svc.clock.Now().Add(5 * time.Second)) {
// Give the client 5 seconds to connect before considering the
// campaign orphaned
return osqueryError{message: "campaign waiting for listener"}
}
+2 -2
View File
@@ -36,8 +36,8 @@ func decodeListHostsRequest(ctx context.Context, r *http.Request) (interface{},
opt, err := listOptionsFromRequest(r)
hopt := kolide.HostListOptions{ListOptions: opt}
status := r.URL.Query().Get("status")
switch status {
case "new", "online", "offline", "mia":
switch kolide.HostStatus(status) {
case kolide.StatusNew, kolide.StatusOnline, kolide.StatusOffline, kolide.StatusMIA:
hopt.StatusFilter = kolide.HostStatus(status)
case "":
// No error when unset