source: server/db/db.go@ fe28254

main
Last change on this file since fe28254 was b715712, checked in by Stefan <trsunovstefan@…>, 8 weeks ago

Add the server side and configuration

  • Property mode set to 100644
File size: 3.6 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
13 _ "github.com/lib/pq"
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//
21//go:embed schema_creation.sql data_load.sql
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
38 var err error
39 DB, err = sql.Open("postgres", dsn)
40 if err != nil {
41 return fmt.Errorf("sql.Open: %w", err)
42 }
43 if err := DB.Ping(); err != nil {
44 return fmt.Errorf("db ping (host=%s port=%s user=%s dbname=%s): %w",
45 host, port, user, name, err)
46 }
47 return nil
48}
49
50// runScript executes one of the embedded .sql scripts as a single statement.
51func runScript(name string) error {
52 content, err := sqlScripts.ReadFile(name)
53 if err != nil {
54 return fmt.Errorf("read embedded %s: %w", name, err)
55 }
56 if _, err := DB.Exec(string(content)); err != nil {
57 return fmt.Errorf("exec %s: %w", name, err)
58 }
59 return nil
60}
61
62// RunSQLFile executes a .sql file from disk as a single statement.
63func RunSQLFile(path string) error {
64 content, err := os.ReadFile(path)
65 if err != nil {
66 return fmt.Errorf("read %s: %w", path, err)
67 }
68 if _, err := DB.Exec(string(content)); err != nil {
69 return fmt.Errorf("exec %s: %w", path, err)
70 }
71 return nil
72}
73
74// InitSchema runs schema_creation.sql then data_load.sql.
75// Destructive: drops the `project` schema. Intended for the -init flag.
76func InitSchema() error {
77 log.Println("Running schema_creation.sql ...")
78 if err := runScript("schema_creation.sql"); err != nil {
79 return err
80 }
81 if err := LoadData(); err != nil {
82 return err
83 }
84 log.Println("Database initialised.")
85 return nil
86}
87
88// LoadData reloads the sample data without touching the schema.
89func LoadData() error {
90 log.Println("Running data_load.sql ...")
91 return runScript("data_load.sql")
92}
93
94// loadEnvFile looks for a .env file in the working directory and in every
95// parent directory, so the program can be started from the repo root, from
96// server/, or from anywhere else inside the checkout. Variables already set
97// in the real environment always win over the file, which is what lets you
98// point the prototype at the faculty database with DBHOST=... ./eduberza
99func loadEnvFile() {
100 dir, err := os.Getwd()
101 if err != nil {
102 return
103 }
104 for {
105 path := filepath.Join(dir, ".env")
106 if applyEnvFile(path) {
107 return
108 }
109 parent := filepath.Dir(dir)
110 if parent == dir {
111 return // reached the filesystem root
112 }
113 dir = parent
114 }
115}
116
117// applyEnvFile reports whether the file existed and was read.
118func applyEnvFile(path string) bool {
119 f, err := os.Open(path)
120 if err != nil {
121 return false
122 }
123 defer f.Close()
124
125 s := bufio.NewScanner(f)
126 for s.Scan() {
127 line := strings.TrimSpace(s.Text())
128 if line == "" || strings.HasPrefix(line, "#") {
129 continue
130 }
131 key, value, ok := strings.Cut(line, "=")
132 if !ok {
133 continue
134 }
135 key = strings.TrimSpace(key)
136 // Do not clobber variables that are already set in the environment.
137 if _, exists := os.LookupEnv(key); exists {
138 continue
139 }
140 os.Setenv(key, strings.Trim(strings.TrimSpace(value), `"'`))
141 }
142 if err := s.Err(); err != nil {
143 log.Printf("warning: could not fully read %s: %v", path, err)
144 }
145 return true
146}
147
148func getenv(key, def string) string {
149 if v := os.Getenv(key); v != "" {
150 return v
151 }
152 return def
153}
Note: See TracBrowser for help on using the repository browser.