feat: introduce database type (#982)

See https://github.com/ooni/probe/issues/2352

Co-authored-by: decfox <decfox@github.com>
This commit is contained in:
DecFox 2022-11-16 20:21:41 +05:30 committed by GitHub
commit 28aabe0947
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
14 changed files with 144 additions and 94 deletions

View file

@ -17,10 +17,31 @@ import (
"github.com/upper/db/v4"
)
// Open returns a new database instance
func Open(dbpath string) (*Database, error) {
db, err := Connect(dbpath)
if err != nil {
return nil, err
}
return &Database{
sess: db,
}, nil
}
// Database is a database instance to store measurements
type Database struct {
sess db.Session
}
// Session returns the database session
func (d *Database) Session() db.Session {
return d.sess
}
// ListMeasurements given a result ID
func ListMeasurements(sess db.Session, resultID int64) ([]MeasurementURLNetwork, error) {
func (d *Database) ListMeasurements(resultID int64) ([]MeasurementURLNetwork, error) {
measurements := []MeasurementURLNetwork{}
req := sess.SQL().Select(
req := d.sess.SQL().Select(
db.Raw("networks.*"),
db.Raw("urls.*"),
db.Raw("measurements.*"),
@ -39,12 +60,12 @@ func ListMeasurements(sess db.Session, resultID int64) ([]MeasurementURLNetwork,
}
// GetMeasurementJSON returns a map[string]interface{} given a database and a measurementID
func GetMeasurementJSON(sess db.Session, measurementID int64) (map[string]interface{}, error) {
func (d *Database) GetMeasurementJSON(measurementID int64) (map[string]interface{}, error) {
var (
measurement MeasurementURLNetwork
msmtJSON map[string]interface{}
)
req := sess.SQL().Select(
req := d.sess.SQL().Select(
db.Raw("urls.*"),
db.Raw("measurements.*"),
).From("measurements").
@ -102,10 +123,10 @@ func GetMeasurementJSON(sess db.Session, measurementID int64) (map[string]interf
}
// ListResults return the list of results
func ListResults(sess db.Session) ([]ResultNetwork, []ResultNetwork, error) {
func (d *Database) ListResults() ([]ResultNetwork, []ResultNetwork, error) {
doneResults := []ResultNetwork{}
incompleteResults := []ResultNetwork{}
req := sess.SQL().Select(
req := d.sess.SQL().Select(
db.Raw("networks.network_name"),
db.Raw("networks.network_type"),
db.Raw("networks.ip"),
@ -165,9 +186,9 @@ func ListResults(sess db.Session) ([]ResultNetwork, []ResultNetwork, error) {
// DeleteResult will delete a particular result and the relative measurement on
// disk.
func DeleteResult(sess db.Session, resultID int64) error {
func (d *Database) DeleteResult(resultID int64) error {
var result Result
res := sess.Collection("results").Find("result_id", resultID)
res := d.sess.Collection("results").Find("result_id", resultID)
if err := res.One(&result); err != nil {
if err == db.ErrNoMoreRows {
return err
@ -185,8 +206,8 @@ func DeleteResult(sess db.Session, resultID int64) error {
}
// UpdateUploadedStatus will check if all the measurements inside of a given result set have been uploaded and if so will set the is_uploaded flag to true
func UpdateUploadedStatus(sess db.Session, result *Result) error {
err := sess.Tx(func(tx db.Session) error {
func (d *Database) UpdateUploadedStatus(result *Result) error {
err := d.sess.Tx(func(tx db.Session) error {
uploadedTotal := UploadedTotalCount{}
req := tx.SQL().Select(
db.Raw("SUM(measurements.measurement_is_uploaded)"),
@ -222,7 +243,7 @@ func UpdateUploadedStatus(sess db.Session, result *Result) error {
// CreateMeasurement writes the measurement to the database a returns a pointer
// to the Measurement
func CreateMeasurement(sess db.Session, reportID sql.NullString, testName string, measurementDir string, idx int, resultID int64, urlID sql.NullInt64) (*Measurement, error) {
func (d *Database) CreateMeasurement(reportID sql.NullString, testName string, measurementDir string, idx int, resultID int64, urlID sql.NullInt64) (*Measurement, error) {
// TODO we should look into generating this file path in a more robust way.
// If there are two identical test_names in the same test group there is
// going to be a clash of test_name
@ -240,7 +261,7 @@ func CreateMeasurement(sess db.Session, reportID sql.NullString, testName string
TestKeys: "",
}
newID, err := sess.Collection("measurements").Insert(msmt)
newID, err := d.sess.Collection("measurements").Insert(msmt)
if err != nil {
return nil, errors.Wrap(err, "creating measurement")
}
@ -250,7 +271,7 @@ func CreateMeasurement(sess db.Session, reportID sql.NullString, testName string
// CreateResult writes the Result to the database a returns a pointer
// to the Result
func CreateResult(sess db.Session, homePath string, testGroupName string, networkID int64) (*Result, error) {
func (d *Database) CreateResult(homePath string, testGroupName string, networkID int64) (*Result, error) {
startTime := time.Now().UTC()
p, err := makeResultsDir(homePath, testGroupName, startTime)
@ -266,7 +287,7 @@ func CreateResult(sess db.Session, homePath string, testGroupName string, networ
result.MeasurementDir = p
log.Debugf("Creating result %v", result)
newID, err := sess.Collection("results").Insert(result)
newID, err := d.sess.Collection("results").Insert(result)
if err != nil {
return nil, errors.Wrap(err, "creating result")
}
@ -275,7 +296,7 @@ func CreateResult(sess db.Session, homePath string, testGroupName string, networ
}
// CreateNetwork will create a new network in the network table
func CreateNetwork(sess db.Session, loc engine.LocationProvider) (*Network, error) {
func (d *Database) CreateNetwork(loc engine.LocationProvider) (*Network, error) {
network := Network{
ASN: loc.ProbeASN(),
CountryCode: loc.ProbeCC(),
@ -284,7 +305,7 @@ func CreateNetwork(sess db.Session, loc engine.LocationProvider) (*Network, erro
NetworkType: "wifi",
IP: loc.ProbeIP(),
}
newID, err := sess.Collection("networks").Insert(network)
newID, err := d.sess.Collection("networks").Insert(network)
if err != nil {
return nil, err
}
@ -296,10 +317,10 @@ func CreateNetwork(sess db.Session, loc engine.LocationProvider) (*Network, erro
// CreateOrUpdateURL will create a new URL entry to the urls table if it doesn't
// exists, otherwise it will update the category code of the one already in
// there.
func CreateOrUpdateURL(sess db.Session, urlStr string, categoryCode string, countryCode string) (int64, error) {
func (d *Database) CreateOrUpdateURL(urlStr string, categoryCode string, countryCode string) (int64, error) {
var url URL
err := sess.Tx(func(tx db.Session) error {
err := d.sess.Tx(func(tx db.Session) error {
res := tx.Collection("urls").Find(
db.Cond{"url": urlStr, "url_country_code": countryCode},
)
@ -336,7 +357,7 @@ func CreateOrUpdateURL(sess db.Session, urlStr string, categoryCode string, coun
}
// AddTestKeys writes the summary to the measurement
func AddTestKeys(sess db.Session, msmt *Measurement, tk interface{}) error {
func (d *Database) AddTestKeys(msmt *Measurement, tk interface{}) error {
var (
isAnomaly bool
isAnomalyValid bool
@ -357,10 +378,15 @@ func AddTestKeys(sess db.Session, msmt *Measurement, tk interface{}) error {
msmt.TestKeys = string(tkBytes)
msmt.IsAnomaly = sql.NullBool{Bool: isAnomaly, Valid: isAnomalyValid}
err = sess.Collection("measurements").Find("measurement_id", msmt.ID).Update(msmt)
err = d.sess.Collection("measurements").Find("measurement_id", msmt.ID).Update(msmt)
if err != nil {
log.WithError(err).Error("failed to update measurement")
return errors.Wrap(err, "updating measurement")
}
return nil
}
// Close closes the database session
func (d *Database) Close() error {
return d.sess.Close()
}

View file

@ -43,6 +43,31 @@ func (lp *locationInfo) ResolverIP() string {
return lp.resolverIP
}
func TestNewDatabase(t *testing.T) {
t.Run("with empty path", func(t *testing.T) {
dbpath := ""
database, err := Open(dbpath)
if database != nil {
t.Fatal("unexpected database instance")
}
if err.Error() != "Expecting file:// connection scheme." {
t.Fatal(err)
}
})
t.Run("with valid path", func(t *testing.T) {
tmpfile, err := ioutil.TempFile("", "dbtest")
if err != nil {
t.Fatal(err)
}
defer os.Remove(tmpfile.Name())
_, err = Open(tmpfile.Name())
if err != nil {
t.Fatal(err)
}
})
}
func TestMeasurementWorkflow(t *testing.T) {
tmpfile, err := ioutil.TempFile("", "dbtest")
if err != nil {
@ -56,7 +81,7 @@ func TestMeasurementWorkflow(t *testing.T) {
}
defer os.RemoveAll(tmpdir)
sess, err := Connect(tmpfile.Name())
database, err := Open(tmpfile.Name())
if err != nil {
t.Fatal(err)
}
@ -66,12 +91,13 @@ func TestMeasurementWorkflow(t *testing.T) {
countryCode: "IT",
networkName: "Unknown",
}
network, err := CreateNetwork(sess, &location)
sess := database.Session()
network, err := database.CreateNetwork(&location)
if err != nil {
t.Fatal(err)
}
result, err := CreateResult(sess, tmpdir, "websites", network.ID)
result, err := database.CreateResult(tmpdir, "websites", network.ID)
if err != nil {
t.Fatal(err)
}
@ -82,7 +108,7 @@ func TestMeasurementWorkflow(t *testing.T) {
msmtFilePath := tmpdir
urlID := sql.NullInt64{Int64: 0, Valid: false}
m1, err := CreateMeasurement(sess, reportID, testName, msmtFilePath, 0, resultID, urlID)
m1, err := database.CreateMeasurement(reportID, testName, msmtFilePath, 0, resultID, urlID)
if err != nil {
t.Fatal(err)
}
@ -93,7 +119,7 @@ func TestMeasurementWorkflow(t *testing.T) {
t.Fatal(err)
}
m2, err := CreateMeasurement(sess, reportID, testName, msmtFilePath, 0, resultID, urlID)
m2, err := database.CreateMeasurement(reportID, testName, msmtFilePath, 0, resultID, urlID)
if err != nil {
t.Fatal(err)
}
@ -107,7 +133,7 @@ func TestMeasurementWorkflow(t *testing.T) {
if m2.ResultID != m1.ResultID {
t.Error("result_id mismatch")
}
err = UpdateUploadedStatus(sess, result)
err = database.UpdateUploadedStatus(result)
if err != nil {
t.Fatal(err)
}
@ -122,7 +148,7 @@ func TestMeasurementWorkflow(t *testing.T) {
t.Error("result should be marked as not uploaded")
}
done, incomplete, err := ListResults(sess)
done, incomplete, err := database.ListResults()
if err != nil {
t.Fatal(err)
}
@ -141,7 +167,7 @@ func TestMeasurementWorkflow(t *testing.T) {
t.Error("there should be a total of 1 anomalies in the result")
}
msmts, err := ListMeasurements(sess, resultID)
msmts, err := database.ListMeasurements(resultID)
if err != nil {
t.Fatal(err)
}
@ -163,7 +189,7 @@ func TestDeleteResult(t *testing.T) {
}
defer os.RemoveAll(tmpdir)
sess, err := Connect(tmpfile.Name())
database, err := Open(tmpfile.Name())
if err != nil {
t.Fatal(err)
}
@ -173,12 +199,13 @@ func TestDeleteResult(t *testing.T) {
countryCode: "IT",
networkName: "Unknown",
}
network, err := CreateNetwork(sess, &location)
sess := database.Session()
network, err := database.CreateNetwork(&location)
if err != nil {
t.Fatal(err)
}
result, err := CreateResult(sess, tmpdir, "websites", network.ID)
result, err := database.CreateResult(tmpdir, "websites", network.ID)
if err != nil {
t.Fatal(err)
}
@ -189,7 +216,7 @@ func TestDeleteResult(t *testing.T) {
msmtFilePath := tmpdir
urlID := sql.NullInt64{Int64: 0, Valid: false}
m1, err := CreateMeasurement(sess, reportID, testName, msmtFilePath, 0, resultID, urlID)
m1, err := database.CreateMeasurement(reportID, testName, msmtFilePath, 0, resultID, urlID)
if err != nil {
t.Fatal(err)
}
@ -204,7 +231,7 @@ func TestDeleteResult(t *testing.T) {
t.Error("result_id mismatch")
}
err = DeleteResult(sess, resultID)
err = database.DeleteResult(resultID)
if err != nil {
t.Fatal(err)
}
@ -223,7 +250,7 @@ func TestDeleteResult(t *testing.T) {
t.Fatal("measurements should be zero")
}
err = DeleteResult(sess, 20)
err = database.DeleteResult(20)
if err != db.ErrNoMoreRows {
t.Fatal(err)
}
@ -236,7 +263,7 @@ func TestNetworkCreate(t *testing.T) {
}
defer os.Remove(tmpfile.Name())
sess, err := Connect(tmpfile.Name())
database, err := Open(tmpfile.Name())
if err != nil {
t.Fatal(err)
}
@ -253,12 +280,12 @@ func TestNetworkCreate(t *testing.T) {
networkName: "Fufnet",
}
_, err = CreateNetwork(sess, &l1)
_, err = database.CreateNetwork(&l1)
if err != nil {
t.Fatal(err)
}
_, err = CreateNetwork(sess, &l2)
_, err = database.CreateNetwork(&l2)
if err != nil {
t.Fatal(err)
}
@ -272,32 +299,32 @@ func TestURLCreation(t *testing.T) {
}
defer os.Remove(tmpfile.Name())
sess, err := Connect(tmpfile.Name())
database, err := Open(tmpfile.Name())
if err != nil {
t.Fatal(err)
}
newID1, err := CreateOrUpdateURL(sess, "https://google.com", "GMB", "XX")
newID1, err := database.CreateOrUpdateURL("https://google.com", "GMB", "XX")
if err != nil {
t.Fatal(err)
}
newID2, err := CreateOrUpdateURL(sess, "https://google.com", "SRCH", "XX")
newID2, err := database.CreateOrUpdateURL("https://google.com", "SRCH", "XX")
if err != nil {
t.Fatal(err)
}
newID3, err := CreateOrUpdateURL(sess, "https://facebook.com", "GRP", "XX")
newID3, err := database.CreateOrUpdateURL("https://facebook.com", "GRP", "XX")
if err != nil {
t.Fatal(err)
}
newID4, err := CreateOrUpdateURL(sess, "https://facebook.com", "GMP", "XX")
newID4, err := database.CreateOrUpdateURL("https://facebook.com", "GMP", "XX")
if err != nil {
t.Fatal(err)
}
newID5, err := CreateOrUpdateURL(sess, "https://google.com", "SRCH", "XX")
newID5, err := database.CreateOrUpdateURL("https://google.com", "SRCH", "XX")
if err != nil {
t.Fatal(err)
}
@ -351,22 +378,22 @@ func TestGetMeasurementJSON(t *testing.T) {
}
defer os.RemoveAll(tmpdir)
sess, err := Connect(tmpfile.Name())
database, err := Open(tmpfile.Name())
if err != nil {
t.Fatal(err)
}
sess := database.Session()
location := locationInfo{
asn: 0,
countryCode: "IT",
networkName: "Unknown",
}
network, err := CreateNetwork(sess, &location)
network, err := database.CreateNetwork(&location)
if err != nil {
t.Fatal(err)
}
result, err := CreateResult(sess, tmpdir, "websites", network.ID)
result, err := database.CreateResult(tmpdir, "websites", network.ID)
if err != nil {
t.Fatal(err)
}
@ -377,7 +404,7 @@ func TestGetMeasurementJSON(t *testing.T) {
msmtFilePath := tmpdir
urlID := sql.NullInt64{Int64: 0, Valid: false}
msmt, err := CreateMeasurement(sess, reportID, testName, msmtFilePath, 0, resultID, urlID)
msmt, err := database.CreateMeasurement(reportID, testName, msmtFilePath, 0, resultID, urlID)
if err != nil {
t.Fatal(err)
}
@ -387,7 +414,7 @@ func TestGetMeasurementJSON(t *testing.T) {
t.Fatal(err)
}
tk, err := GetMeasurementJSON(sess, msmt.ID)
tk, err := database.GetMeasurementJSON(msmt.ID)
if err != nil {
t.Fatal(err)
}