source: server/db/db.go@ ef1c1c7

main
Last change on this file since ef1c1c7 was ef1c1c7, checked in by Stefan <trsunovstefan@…>, 6 days ago

Wiki docs, phase 6 and phase 7 added

  • Property mode set to 100644
File size: 4.2 KB
RevLine 
[b715712]1package db
2
3import (
4 "bufio"
5 "database/sql"
6 "embed"
7 "fmt"
8 "log"
9 "os"
10 "path/filepath"
11 "strings"
12
[ef1c1c7]13 "github.com/lib/pq"
[b715712]14)
15
16var DB *sql.DB
17
18// The SQL scripts are compiled into the binary so that -init works no matter
19// which directory the program is started from.
20//
[ef1c1c7]21//go:embed schema_creation.sql advanced_db.sql data_load.sql
[b715712]22var sqlScripts embed.FS
23
24func Connect() error {
25 loadEnvFile()
26
27 host := getenv("DBHOST", "localhost")
28 port := getenv("DBPORT", "5432")
29 user := getenv("DBUSER", "postgres")
30 pass := getenv("DBPASSWORD", "")
31 name := getenv("DBNAME", "postgres")
32
33 dsn := fmt.Sprintf(
34 "host=%s port=%s user=%s password=%s dbname=%s sslmode=disable options='--search_path=project,public'",
35 host, port, user, pass, name,
36 )
37
[ef1c1c7]38 // Optional SSH tunnel, the same thing DBeaver's "SSH" tab does. When
39 // SSH_HOST is set, DBHOST/DBPORT are resolved from the SSH server's side
40 // (for the faculty server that is usually localhost:5432).
41 if sshHost := os.Getenv("SSH_HOST"); sshHost != "" {
42 dialer, err := newSSHDialer(sshHost)
43 if err != nil {
44 return err
45 }
46 connector, err := pq.NewConnector(dsn)
47 if err != nil {
48 return fmt.Errorf("pq.NewConnector: %w", err)
49 }
50 connector.Dialer(dialer)
51 DB = sql.OpenDB(connector)
52 } else {
53 var err error
54 DB, err = sql.Open("postgres", dsn)
55 if err != nil {
56 return fmt.Errorf("sql.Open: %w", err)
57 }
[b715712]58 }
59 if err := DB.Ping(); err != nil {
60 return fmt.Errorf("db ping (host=%s port=%s user=%s dbname=%s): %w",
61 host, port, user, name, err)
62 }
63 return nil
64}
65
66// runScript executes one of the embedded .sql scripts as a single statement.
67func runScript(name string) error {
68 content, err := sqlScripts.ReadFile(name)
69 if err != nil {
70 return fmt.Errorf("read embedded %s: %w", name, err)
71 }
72 if _, err := DB.Exec(string(content)); err != nil {
73 return fmt.Errorf("exec %s: %w", name, err)
74 }
75 return nil
76}
77
78// RunSQLFile executes a .sql file from disk as a single statement.
79func RunSQLFile(path string) error {
80 content, err := os.ReadFile(path)
81 if err != nil {
82 return fmt.Errorf("read %s: %w", path, err)
83 }
84 if _, err := DB.Exec(string(content)); err != nil {
85 return fmt.Errorf("exec %s: %w", path, err)
86 }
87 return nil
88}
89
[ef1c1c7]90// InitSchema runs schema_creation.sql, advanced_db.sql (P7) and data_load.sql.
[b715712]91// Destructive: drops the `project` schema. Intended for the -init flag.
92func InitSchema() error {
[ef1c1c7]93 for _, name := range []string{"schema_creation.sql", "advanced_db.sql"} {
94 log.Printf("Running %s ...", name)
95 if err := runScript(name); err != nil {
96 return err
97 }
[b715712]98 }
99 if err := LoadData(); err != nil {
100 return err
101 }
102 log.Println("Database initialised.")
103 return nil
104}
105
106// LoadData reloads the sample data without touching the schema.
107func LoadData() error {
108 log.Println("Running data_load.sql ...")
109 return runScript("data_load.sql")
110}
111
112// loadEnvFile looks for a .env file in the working directory and in every
113// parent directory, so the program can be started from the repo root, from
114// server/, or from anywhere else inside the checkout. Variables already set
115// in the real environment always win over the file, which is what lets you
116// point the prototype at the faculty database with DBHOST=... ./eduberza
117func loadEnvFile() {
118 dir, err := os.Getwd()
119 if err != nil {
120 return
121 }
122 for {
123 path := filepath.Join(dir, ".env")
124 if applyEnvFile(path) {
125 return
126 }
127 parent := filepath.Dir(dir)
128 if parent == dir {
129 return // reached the filesystem root
130 }
131 dir = parent
132 }
133}
134
135// applyEnvFile reports whether the file existed and was read.
136func applyEnvFile(path string) bool {
137 f, err := os.Open(path)
138 if err != nil {
139 return false
140 }
141 defer f.Close()
142
143 s := bufio.NewScanner(f)
144 for s.Scan() {
145 line := strings.TrimSpace(s.Text())
146 if line == "" || strings.HasPrefix(line, "#") {
147 continue
148 }
149 key, value, ok := strings.Cut(line, "=")
150 if !ok {
151 continue
152 }
153 key = strings.TrimSpace(key)
154 // Do not clobber variables that are already set in the environment.
155 if _, exists := os.LookupEnv(key); exists {
156 continue
157 }
158 os.Setenv(key, strings.Trim(strings.TrimSpace(value), `"'`))
159 }
160 if err := s.Err(); err != nil {
161 log.Printf("warning: could not fully read %s: %v", path, err)
162 }
163 return true
164}
165
166func getenv(key, def string) string {
167 if v := os.Getenv(key); v != "" {
168 return v
169 }
170 return def
171}
Note: See TracBrowser for help on using the repository browser.