Skip to content

Commit b169fcd

Browse files
authored
Merge pull request #4 from open-oni/feature/ensure-awardee
Feature/ensure awardee
2 parents 8829b3f + b1ad058 commit b169fcd

7 files changed

Lines changed: 146 additions & 36 deletions

File tree

.gitattributes

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
*.go diff=golang

README.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,7 @@ export BA_BIND=":2222"
5050
export BATCH_SOURCE="/mnt/news/production-batches"
5151
export ONI_LOCATION="/opt/openoni/"
5252
export HOST_KEY_FILE="/etc/oni-agent"
53+
export DB_CONNECTION="user:password@tcp(127.0.0.1:3306)/databasename"
5354
make
5455
./bin/agent
5556
```

cmd/agent/main.go

Lines changed: 57 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ import (
55
"crypto/rand"
66
"crypto/rsa"
77
"crypto/x509"
8+
"database/sql"
89
"encoding/pem"
910
"errors"
1011
"fmt"
@@ -13,6 +14,7 @@ import (
1314
"sync/atomic"
1415

1516
gliderssh "github.com/gliderlabs/ssh"
17+
_ "github.com/go-sql-driver/mysql"
1618
"github.com/open-oni/oni-agent/internal/queue"
1719
"github.com/open-oni/oni-agent/internal/version"
1820
"golang.org/x/crypto/ssh"
@@ -35,48 +37,78 @@ var HostKeySigner ssh.Signer
3537
// background jobs, providing status of existing jobs, etc.
3638
var JobRunner *queue.Queue
3739

40+
// dbPool is our single DB connection shared app-wide
41+
var dbPool *sql.DB
42+
3843
func getEnvironment() {
44+
var errList []error
45+
var err error
46+
3947
BABind = os.Getenv("BA_BIND")
4048
if BABind == "" {
41-
slog.Error("BA_BIND must be set")
42-
os.Exit(1)
49+
errList = append(errList, errors.New("BA_BIND must be set"))
4350
}
4451

4552
ONILocation = os.Getenv("ONI_LOCATION")
46-
47-
var info, err = os.Stat(ONILocation)
48-
if err == nil {
49-
if !info.IsDir() {
50-
err = errors.New("not a valid directory")
53+
if ONILocation == "" {
54+
errList = append(errList, errors.New("ONI_LOCATION must be set"))
55+
} else {
56+
var info, err = os.Stat(ONILocation)
57+
if err == nil {
58+
if !info.IsDir() {
59+
err = errors.New("not a valid directory")
60+
}
61+
}
62+
if err != nil {
63+
errList = append(errList, fmt.Errorf("Invalid setting for ONI_LOCATION: %w", err))
5164
}
52-
}
53-
if err != nil {
54-
slog.Error("Invalid setting for ONI_LOCATION", "error", err)
55-
os.Exit(1)
5665
}
5766

5867
BatchSource = os.Getenv("BATCH_SOURCE")
59-
info, err = os.Stat(BatchSource)
60-
if err == nil {
61-
if !info.IsDir() {
62-
err = errors.New("not a valid directory")
68+
if BatchSource == "" {
69+
errList = append(errList, errors.New("BATCH_SOURCE must be set"))
70+
} else {
71+
var info, err = os.Stat(BatchSource)
72+
if err == nil {
73+
if !info.IsDir() {
74+
err = errors.New("not a valid directory")
75+
}
76+
}
77+
if err != nil {
78+
errList = append(errList, fmt.Errorf("Invalid setting for BATCH_SOURCE: %w", err))
6379
}
64-
}
65-
if err != nil {
66-
slog.Error("Invalid setting for BATCH_SOURCE", "error", err)
67-
os.Exit(1)
6880
}
6981

7082
var fname = os.Getenv("HOST_KEY_FILE")
7183
if fname == "" {
72-
slog.Error("HOST_KEY_FILE must be set")
73-
os.Exit(1)
84+
errList = append(errList, errors.New("HOST_KEY_FILE must be set"))
85+
} else {
86+
HostKeySigner, err = readKey(fname)
87+
if err != nil {
88+
errList = append(errList, fmt.Errorf("HOST_KEY_FILE is invalid or cannot be read: %w", err))
89+
}
7490
}
75-
HostKeySigner, err = readKey(fname)
76-
if err != nil {
77-
slog.Error("HOST_KEY_FILE is invalid or cannot be read", "error", err)
91+
92+
var connect = os.Getenv("DB_CONNECTION")
93+
if connect == "" {
94+
errList = append(errList, errors.New(`DB_CONNECTION must be set (e.g., "user:pass@tcp(127.0.0.1:3306)/dbname")`))
95+
} else {
96+
dbPool, err = sql.Open("mysql", connect)
97+
if err != nil {
98+
errList = append(errList, fmt.Errorf(`DB_CONNECTION is invalid: %w`, err))
99+
}
100+
}
101+
102+
if len(errList) > 0 {
103+
for _, err := range errList {
104+
fmt.Fprintf(os.Stderr, " - %s\n", err)
105+
}
78106
os.Exit(1)
79107
}
108+
109+
dbPool.SetConnMaxLifetime(0)
110+
dbPool.SetMaxIdleConns(3)
111+
dbPool.SetMaxOpenConns(3)
80112
}
81113

82114
func readKey(keyfile string) (ssh.Signer, error) {
@@ -155,6 +187,7 @@ func main() {
155187
trapIntTerm(func() {
156188
cancel()
157189
srv.Close()
190+
dbPool.Close()
158191
})
159192
go JobRunner.Wait(ctx)
160193

cmd/agent/session.go

Lines changed: 79 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
package main
22

33
import (
4+
"database/sql"
45
"encoding/json"
56
"fmt"
67
"log/slog"
@@ -62,44 +63,55 @@ func (s session) respond(st Status, msg string, data H) {
6263
}
6364

6465
func (s session) handle() {
65-
var cmds = s.Command()
66-
if len(cmds) == 0 {
66+
var parts = s.Command()
67+
if len(parts) == 0 {
6768
s.respond(StatusError, "no command specified", nil)
6869
return
6970
}
7071

71-
var command = cmds[0]
72+
var command, args = parts[0], parts[1:]
7273
switch command {
7374
case "version":
7475
s.respond(StatusSuccess, "", H{"version": version.Version})
7576

7677
case "job-status":
77-
if len(cmds) != 2 {
78+
if len(args) != 1 {
7879
s.respond(StatusError, "You must supply a job ID", nil)
7980
return
8081
}
81-
s.getJobStatus(cmds[1])
82+
s.getJobStatus(args[0])
8283

8384
case "job-logs":
84-
if len(cmds) != 2 {
85+
if len(args) != 1 {
8586
s.respond(StatusError, "You must supply a job ID", nil)
8687
return
8788
}
88-
s.getJobLogs(cmds[1])
89+
s.getJobLogs(args[0])
8990

9091
case "load-batch":
91-
if len(cmds) != 2 {
92+
if len(args) != 1 {
9293
s.respond(StatusError, fmt.Sprintf("%q requires exactly one batch name", command), nil)
9394
return
9495
}
95-
s.loadBatch(cmds[1])
96+
s.loadBatch(args[0])
9697

9798
case "purge-batch":
98-
if len(cmds) != 2 {
99+
if len(args) != 1 {
99100
s.respond(StatusError, fmt.Sprintf("%q requires exactly one batch name", command), nil)
100101
return
101102
}
102-
s.purgeBatch(cmds[1])
103+
s.purgeBatch(args[0])
104+
105+
case "ensure-awardee":
106+
if len(args) < 1 || len(args) > 2 {
107+
s.respond(StatusError, fmt.Sprintf("%q requires one or two args: MARC org code and awardee name. Name is required if the awardee is to be auto-created.", command), nil)
108+
return
109+
}
110+
111+
if len(args) == 1 {
112+
args = []string{args[0], ""}
113+
}
114+
s.ensureAwardee(args[0], args[1])
103115

104116
default:
105117
s.respond(StatusError, fmt.Sprintf("%q is not a valid command name", command), nil)
@@ -185,6 +197,62 @@ func (s session) queueJob(command string, args ...string) {
185197
s.respond(StatusSuccess, "Job added to queue", H{"job": H{"id": id}})
186198
}
187199

200+
func (s session) ensureAwardee(code string, name string) {
201+
var rows, err = dbPool.Query("SELECT COUNT(*) FROM core_awardee WHERE org_code = ?", code)
202+
if err != nil {
203+
s.respond(StatusError, "Unable to query database", H{"error": err.Error()})
204+
return
205+
}
206+
defer rows.Close()
207+
208+
// What does it mean if there's no error reported, but no count returned?
209+
if !rows.Next() {
210+
s.respond(StatusError, "Unable to count awardees in database", H{"error": "no rows returned by SQL COUNT()"})
211+
return
212+
}
213+
214+
var count int
215+
err = rows.Scan(&count)
216+
if err != nil {
217+
s.respond(StatusError, "Unable to count awardees in database", H{"error": err.Error()})
218+
return
219+
}
220+
221+
// We really only care that there's at least one row. If there are dupes,
222+
// that's out of scope to deal with, and technically not an error in terms of
223+
// what we need.
224+
if count > 0 {
225+
s.respond(StatusSuccess, "Awardee already exists", nil)
226+
return
227+
}
228+
229+
// No rows, no error: if a name was given, create the awardee, otherwise abort
230+
if name == "" {
231+
s.respond(StatusError, "Unable to create awardee", H{"error": "awardee name must be given to auto-create awardees", "org_code": code, "name": name})
232+
return
233+
}
234+
235+
var result sql.Result
236+
result, err = dbPool.Exec("INSERT INTO core_awardee (`org_code`, `name`, `created`) VALUES(?, ?, NOW())", code, name)
237+
if err != nil {
238+
s.respond(StatusError, "Unable to create awardee", H{"error": err.Error(), "org_code": code, "name": name})
239+
return
240+
}
241+
var n int64
242+
n, err = result.RowsAffected()
243+
if err != nil {
244+
s.respond(StatusError, "Unable to read result of INSERT", H{"error": err.Error(), "org_code": code, "name": name})
245+
return
246+
}
247+
if n != 1 {
248+
s.respond(StatusError, "Unable to create awardee", H{"error": "No rows created", "org_code": code, "name": name})
249+
return
250+
}
251+
252+
s.respond(StatusSuccess, "Awardee created", nil)
253+
return
254+
}
255+
188256
// close terminates the session, always with a status of 0: Go ssh clients
189257
// return an error if the request is anything but successful, so the caller has
190258
// to parse the status instead.

go.mod

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,9 @@ go 1.22.5
55
require github.com/gliderlabs/ssh v0.3.7
66

77
require (
8+
filippo.io/edwards25519 v1.1.0 // indirect
89
github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be // indirect
10+
github.com/go-sql-driver/mysql v1.8.1 // indirect
911
github.com/google/go-cmp v0.6.0 // indirect
1012
golang.org/x/crypto v0.27.0 // indirect
1113
golang.org/x/sys v0.25.0 // indirect

go.sum

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,11 @@
1+
filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA=
2+
filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4=
13
github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be h1:9AeTilPcZAjCFIImctFaOjnTIavg87rW78vTPkQqLI8=
24
github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be/go.mod h1:ySMOLuWl6zY27l47sB3qLNK6tF2fkHG55UZxx8oIVo4=
35
github.com/gliderlabs/ssh v0.3.7 h1:iV3Bqi942d9huXnzEF2Mt+CY9gLu8DNM4Obd+8bODRE=
46
github.com/gliderlabs/ssh v0.3.7/go.mod h1:zpHEXBstFnQYtGnB8k8kQLol82umzn/2/snG7alWVD8=
7+
github.com/go-sql-driver/mysql v1.8.1 h1:LedoTUt/eveggdHS9qUFC1EFSa8bU2+1pZjSRpvNJ1Y=
8+
github.com/go-sql-driver/mysql v1.8.1/go.mod h1:wEBSXgmK//2ZFJyE+qWnIsVGmvmEKlqwuVSjsCm7DZg=
59
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
610
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
711
golang.org/x/crypto v0.27.0 h1:GXm2NjJrPaiv/h1tb2UH8QfgC/hOf/+z0p6PT8o1w7A=

oni-agent.service

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,9 +5,10 @@
55

66
[Service]
77
Environment="BA_BIND=:2222"
8-
Environment="BATCH_SOURCE=/mnt/news/production-batches"
98
Environment="ONI_LOCATION=/opt/openoni/"
9+
Environment="BATCH_SOURCE=/mnt/news/production-batches"
1010
Environment="HOST_KEY_FILE=/etc/oni-agent"
11+
Environment="DB_CONNECTION=user:password@tcp(127.0.0.1:3306)/databasename"
1112
Type=simple
1213
ExecStart=/usr/local/oni-agent/agent
1314
SyslogIdentifier=oni-agent

0 commit comments

Comments
 (0)