Files
team-presence/internal/app/store.go
T
2026-07-28 02:04:36 +03:30

598 lines
17 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package app
import (
"context"
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/base64"
"errors"
"fmt"
"net/mail"
"os"
"strconv"
"strings"
"time"
"unicode"
"unicode/utf8"
_ "modernc.org/sqlite"
)
type User struct {
ID int64
Username string
DisplayName string
Email string
Role string
AvatarURL string
}
type Attendance struct {
ID int64
UserID int64
Day string
CheckIn sql.NullTime
CheckOut sql.NullTime
Mode string
}
type Request struct {
ID int64
UserID int64
UserName string
Kind string
StartDate string
EndDate string
Reason string
Status string
AdminNote string
CreatedAt time.Time
ReviewedAt sql.NullString
Reviewer string
DayCount int
StartJalali string
EndJalali string
}
type DayRosterRow struct {
User User
CheckIn string
CheckOut string
Mode string
RequestKind string
}
type Store struct{ db *sql.DB }
func OpenStore(path string) (*Store, error) {
db, err := sql.Open("sqlite", path+"?_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)&_pragma=foreign_keys(1)")
if err != nil {
return nil, err
}
db.SetMaxOpenConns(1)
s := &Store{db: db}
if err := s.migrate(); err != nil {
db.Close()
return nil, err
}
return s, nil
}
func (s *Store) Close() error { return s.db.Close() }
func (s *Store) Ping(ctx context.Context) error { return s.db.PingContext(ctx) }
func (s *Store) migrate() error {
const schema = `
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT NOT NULL UNIQUE COLLATE NOCASE,
password_hash TEXT,
display_name TEXT NOT NULL,
email TEXT UNIQUE COLLATE NOCASE,
role TEXT NOT NULL DEFAULT 'member' CHECK(role IN ('member','admin')),
avatar_url TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS oauth_accounts (
provider TEXT NOT NULL,
provider_user_id TEXT NOT NULL,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
PRIMARY KEY(provider, provider_user_id)
);
CREATE TABLE IF NOT EXISTS sessions (
token_hash TEXT PRIMARY KEY,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
csrf_token TEXT NOT NULL,
expires_at DATETIME NOT NULL
);
CREATE TABLE IF NOT EXISTS attendance (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
day TEXT NOT NULL,
check_in DATETIME,
check_out DATETIME,
mode TEXT NOT NULL DEFAULT 'office' CHECK(mode IN ('office','remote')),
UNIQUE(user_id, day)
);
CREATE TABLE IF NOT EXISTS requests (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
kind TEXT NOT NULL CHECK(kind IN ('leave','remote')),
start_date TEXT NOT NULL,
end_date TEXT NOT NULL,
reason TEXT NOT NULL DEFAULT '',
status TEXT NOT NULL DEFAULT 'pending' CHECK(status IN ('pending','approved','rejected','cancelled')),
admin_note TEXT NOT NULL DEFAULT '',
reviewed_by INTEGER REFERENCES users(id),
reviewed_at DATETIME,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
CHECK(end_date >= start_date)
);
CREATE INDEX IF NOT EXISTS idx_attendance_day ON attendance(day);
CREATE INDEX IF NOT EXISTS idx_requests_status ON requests(status, start_date);
`
if _, err := s.db.Exec(schema); err != nil {
return err
}
var count int
if err := s.db.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&count); err != nil {
return err
}
if count == 0 {
initialPassword := os.Getenv("INITIAL_ADMIN_PASSWORD")
if initialPassword == "" {
initialPassword = "admin123"
}
hash, err := hashPassword(initialPassword)
if err != nil {
return err
}
_, err = s.db.Exec(`INSERT INTO users(username,password_hash,display_name,email,role) VALUES(?,?,?,?,?)`,
"admin", hash, "Workspace Admin", "admin@localhost", "admin")
return err
}
return nil
}
func hashPassword(password string) (string, error) {
salt := make([]byte, 16)
if _, err := rand.Read(salt); err != nil {
return "", err
}
const rounds = 180000
key := deriveKey([]byte(password), salt, rounds)
return fmt.Sprintf("pbkdf2-sha256$%d$%s$%s", rounds,
base64.RawStdEncoding.EncodeToString(salt), base64.RawStdEncoding.EncodeToString(key)), nil
}
func verifyPassword(encoded, password string) bool {
parts := strings.Split(encoded, "$")
if len(parts) != 4 || parts[0] != "pbkdf2-sha256" {
return false
}
rounds, err := strconv.Atoi(parts[1])
salt, err2 := base64.RawStdEncoding.DecodeString(parts[2])
expected, err3 := base64.RawStdEncoding.DecodeString(parts[3])
if err != nil || err2 != nil || err3 != nil || rounds < 10000 {
return false
}
return hmac.Equal(expected, deriveKey([]byte(password), salt, rounds))
}
func deriveKey(password, salt []byte, rounds int) []byte {
mac := hmac.New(sha256.New, password)
mac.Write(salt)
mac.Write([]byte{0, 0, 0, 1})
u := mac.Sum(nil)
out := append([]byte(nil), u...)
for i := 1; i < rounds; i++ {
mac.Reset()
mac.Write(u)
u = mac.Sum(nil)
for j := range out {
out[j] ^= u[j]
}
}
return out
}
func randomToken(bytes int) (string, error) {
b := make([]byte, bytes)
if _, err := rand.Read(b); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(b), nil
}
func tokenHash(token string) string {
sum := sha256.Sum256([]byte(token))
return base64.RawURLEncoding.EncodeToString(sum[:])
}
func (s *Store) Authenticate(identifier, password string) (*User, error) {
var u User
var hash string
err := s.db.QueryRow(`SELECT id,username,password_hash,display_name,COALESCE(email,''),role,avatar_url
FROM users WHERE username=? OR email=?`, strings.TrimSpace(identifier), strings.TrimSpace(identifier)).
Scan(&u.ID, &u.Username, &hash, &u.DisplayName, &u.Email, &u.Role, &u.AvatarURL)
if err != nil || !verifyPassword(hash, password) {
return nil, errors.New("invalid email, username, or password")
}
return &u, nil
}
func (s *Store) CreateUser(username, password, displayName, email, role string) (int64, error) {
username = strings.TrimSpace(username)
displayName = strings.TrimSpace(displayName)
email = strings.ToLower(strings.TrimSpace(email))
if displayName == "" || len(password) < 8 {
return 0, errors.New("name and a password of at least 8 characters are required")
}
if username == "" && email == "" {
return 0, errors.New("enter an email address or username")
}
if email != "" {
parsed, err := mail.ParseAddress(email)
if err != nil || parsed.Address != email {
return 0, errors.New("enter a valid email address")
}
}
explicitUsername := username != ""
if !explicitUsername {
username = usernameFromEmail(email)
}
if !validUsername(username) {
return 0, errors.New("username must be 340 letters, numbers, dots, dashes, or underscores")
}
if !explicitUsername {
base := username
for suffix := 0; ; suffix++ {
if suffix > 0 {
username = fmt.Sprintf("%s-%d", base, suffix)
}
var count int
if err := s.db.QueryRow(`SELECT COUNT(*) FROM users WHERE username=?`, username).Scan(&count); err != nil {
return 0, err
}
if count == 0 {
break
}
}
}
if role != "admin" {
role = "member"
}
hash, err := hashPassword(password)
if err != nil {
return 0, err
}
result, err := s.db.Exec(`INSERT INTO users(username,password_hash,display_name,email,role)
VALUES(?,?,?,NULLIF(?,''),?)`, username, hash, displayName, email, role)
if err != nil {
if strings.Contains(strings.ToLower(err.Error()), "unique") {
return 0, errors.New("that username or email is already in use")
}
return 0, err
}
return result.LastInsertId()
}
func usernameFromEmail(email string) string {
local := strings.SplitN(email, "@", 2)[0]
var out strings.Builder
for _, r := range local {
if unicode.IsLetter(r) || unicode.IsDigit(r) || strings.ContainsRune("._-", r) {
out.WriteRune(r)
}
}
username := strings.Trim(out.String(), "._-")
if utf8.RuneCountInString(username) < 3 {
username = "member-" + username
}
if utf8.RuneCountInString(username) > 32 {
username = string([]rune(username)[:32])
}
return username
}
func validUsername(username string) bool {
count := utf8.RuneCountInString(username)
if count < 3 || count > 40 {
return false
}
for _, r := range username {
if !unicode.IsLetter(r) && !unicode.IsDigit(r) && !strings.ContainsRune("._-", r) {
return false
}
}
return true
}
func (s *Store) Users() ([]User, error) {
rows, err := s.db.Query(`SELECT id,username,display_name,COALESCE(email,''),role,avatar_url FROM users ORDER BY display_name`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []User
for rows.Next() {
var u User
if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName, &u.Email, &u.Role, &u.AvatarURL); err != nil {
return nil, err
}
out = append(out, u)
}
return out, rows.Err()
}
func (s *Store) CreateSession(userID int64) (token, csrf string, err error) {
token, err = randomToken(32)
if err != nil {
return "", "", err
}
csrf, err = randomToken(24)
if err != nil {
return "", "", err
}
_, err = s.db.Exec(`INSERT INTO sessions(token_hash,user_id,csrf_token,expires_at) VALUES(?,?,?,?)`,
tokenHash(token), userID, csrf, time.Now().Add(30*24*time.Hour))
return
}
func (s *Store) UpsertOAuthUser(provider, providerID, username, email, displayName, avatar string) (int64, error) {
if providerID == "" || providerID == "<nil>" {
return 0, errors.New("OAuth profile did not contain an id")
}
tx, err := s.db.Begin()
if err != nil {
return 0, err
}
defer tx.Rollback()
var userID int64
err = tx.QueryRow(`SELECT user_id FROM oauth_accounts WHERE provider=? AND provider_user_id=?`, provider, providerID).Scan(&userID)
if err == nil {
return userID, tx.Commit()
}
if !errors.Is(err, sql.ErrNoRows) {
return 0, err
}
if email != "" {
err = tx.QueryRow(`SELECT id FROM users WHERE email=?`, email).Scan(&userID)
}
if email == "" || errors.Is(err, sql.ErrNoRows) {
username = strings.TrimSpace(username)
if username == "" {
username = provider + "-" + providerID
}
if displayName == "" {
displayName = username
}
candidate := username
for suffix := 0; ; suffix++ {
if suffix > 0 {
candidate = fmt.Sprintf("%s-%d", username, suffix)
}
var exists int
err = tx.QueryRow(`SELECT COUNT(*) FROM users WHERE username=?`, candidate).Scan(&exists)
if err != nil {
return 0, err
}
if exists == 0 {
break
}
}
result, err := tx.Exec(`INSERT INTO users(username,display_name,email,avatar_url) VALUES(?,?,NULLIF(?,''),?)`,
candidate, displayName, email, avatar)
if err != nil {
return 0, err
}
userID, err = result.LastInsertId()
if err != nil {
return 0, err
}
} else if err != nil {
return 0, err
}
if _, err := tx.Exec(`INSERT INTO oauth_accounts(provider,provider_user_id,user_id) VALUES(?,?,?)`, provider, providerID, userID); err != nil {
return 0, err
}
return userID, tx.Commit()
}
func (s *Store) Session(token string) (*User, string, error) {
var u User
var csrf string
err := s.db.QueryRow(`SELECT u.id,u.username,u.display_name,COALESCE(u.email,''),u.role,u.avatar_url,s.csrf_token
FROM sessions s JOIN users u ON u.id=s.user_id
WHERE s.token_hash=? AND s.expires_at>?`, tokenHash(token), time.Now()).
Scan(&u.ID, &u.Username, &u.DisplayName, &u.Email, &u.Role, &u.AvatarURL, &csrf)
if err != nil {
return nil, "", err
}
return &u, csrf, nil
}
func (s *Store) DeleteSession(token string) {
_, _ = s.db.Exec(`DELETE FROM sessions WHERE token_hash=?`, tokenHash(token))
}
func (s *Store) TodayAttendance(userID int64, day string) (*Attendance, error) {
var a Attendance
err := s.db.QueryRow(`SELECT id,user_id,day,check_in,check_out,mode FROM attendance WHERE user_id=? AND day=?`, userID, day).
Scan(&a.ID, &a.UserID, &a.Day, &a.CheckIn, &a.CheckOut, &a.Mode)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return &a, err
}
func (s *Store) CheckIn(userID int64, day, mode string) error {
if mode != "remote" {
mode = "office"
}
_, err := s.db.Exec(`INSERT INTO attendance(user_id,day,check_in,mode) VALUES(?,?,?,?)
ON CONFLICT(user_id,day) DO UPDATE SET check_in=COALESCE(attendance.check_in,excluded.check_in),mode=excluded.mode`,
userID, day, time.Now(), mode)
return err
}
func (s *Store) CheckOut(userID int64, day string) error {
result, err := s.db.Exec(`UPDATE attendance SET check_out=? WHERE user_id=? AND day=? AND check_in IS NOT NULL AND check_out IS NULL`,
time.Now(), userID, day)
if err != nil {
return err
}
n, _ := result.RowsAffected()
if n == 0 {
return errors.New("check in before checking out")
}
return nil
}
func (s *Store) AttendanceBetween(userID int64, start, end string) (map[string]Attendance, error) {
rows, err := s.db.Query(`SELECT id,user_id,day,check_in,check_out,mode FROM attendance
WHERE user_id=? AND day BETWEEN ? AND ?`, userID, start, end)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[string]Attendance{}
for rows.Next() {
var a Attendance
if err := rows.Scan(&a.ID, &a.UserID, &a.Day, &a.CheckIn, &a.CheckOut, &a.Mode); err != nil {
return nil, err
}
out[a.Day] = a
}
return out, rows.Err()
}
func (s *Store) CreateRequest(userID int64, kind, start, end, reason string) error {
if kind != "leave" && kind != "remote" {
return errors.New("invalid request type")
}
_, err := s.db.Exec(`INSERT INTO requests(user_id,kind,start_date,end_date,reason) VALUES(?,?,?,?,?)`,
userID, kind, start, end, strings.TrimSpace(reason))
return err
}
func (s *Store) Requests(userID int64, admin bool, status string) ([]Request, error) {
query := `SELECT r.id,r.user_id,u.display_name,r.kind,r.start_date,r.end_date,r.reason,r.status,
r.admin_note,r.created_at,r.reviewed_at,COALESCE(a.display_name,'')
FROM requests r JOIN users u ON u.id=r.user_id LEFT JOIN users a ON a.id=r.reviewed_by`
args := []any{}
clauses := []string{}
if !admin {
clauses = append(clauses, "r.user_id=?")
args = append(args, userID)
}
if status != "" {
clauses = append(clauses, "r.status=?")
args = append(args, status)
}
if len(clauses) > 0 {
query += " WHERE " + strings.Join(clauses, " AND ")
}
query += " ORDER BY CASE r.status WHEN 'pending' THEN 0 ELSE 1 END,r.created_at DESC"
rows, err := s.db.Query(query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var out []Request
for rows.Next() {
var r Request
if err := rows.Scan(&r.ID, &r.UserID, &r.UserName, &r.Kind, &r.StartDate, &r.EndDate,
&r.Reason, &r.Status, &r.AdminNote, &r.CreatedAt, &r.ReviewedAt, &r.Reviewer); err != nil {
return nil, err
}
start, _ := time.Parse("2006-01-02", r.StartDate)
end, _ := time.Parse("2006-01-02", r.EndDate)
r.DayCount = int(end.Sub(start).Hours()/24) + 1
out = append(out, r)
}
return out, rows.Err()
}
func (s *Store) ReviewRequest(ctx context.Context, id, reviewerID int64, status, note string) error {
if status != "approved" && status != "rejected" {
return errors.New("invalid review decision")
}
result, err := s.db.ExecContext(ctx, `UPDATE requests SET status=?,admin_note=?,reviewed_by=?,reviewed_at=?
WHERE id=? AND status='pending'`, status, strings.TrimSpace(note), reviewerID, time.Now(), id)
if err != nil {
return err
}
n, _ := result.RowsAffected()
if n == 0 {
return errors.New("request was already reviewed")
}
return nil
}
func (s *Store) CancelRequest(id, userID int64) error {
_, err := s.db.Exec(`UPDATE requests SET status='cancelled' WHERE id=? AND user_id=? AND status='pending'`, id, userID)
return err
}
func (s *Store) ReportRows(start, end string) (*sql.Rows, error) {
return s.db.Query(`SELECT u.display_name,u.username,a.day,
COALESCE(substr(CAST(a.check_in AS TEXT),12,5),''),
COALESCE(substr(CAST(a.check_out AS TEXT),12,5),''),
a.mode,
COALESCE((SELECT r.kind FROM requests r WHERE r.user_id=u.id AND r.status='approved'
AND a.day BETWEEN r.start_date AND r.end_date ORDER BY r.id DESC LIMIT 1),'')
FROM attendance a JOIN users u ON u.id=a.user_id
WHERE a.day BETWEEN ? AND ? ORDER BY a.day,u.display_name`, start, end)
}
func (s *Store) Stats(userID int64, start, end string) (present, remote, leave int, err error) {
err = s.db.QueryRow(`SELECT
COUNT(*),COALESCE(SUM(CASE WHEN mode='remote' THEN 1 ELSE 0 END),0)
FROM attendance WHERE user_id=? AND day BETWEEN ? AND ?`, userID, start, end).Scan(&present, &remote)
if err != nil {
return
}
err = s.db.QueryRow(`SELECT COALESCE(SUM(julianday(end_date)-julianday(start_date)+1),0)
FROM requests WHERE user_id=? AND kind='leave' AND status='approved'
AND start_date<=? AND end_date>=?`, userID, end, start).Scan(&leave)
return
}
func (s *Store) DayRoster(day string) ([]DayRosterRow, error) {
rows, err := s.db.Query(`SELECT
u.id,u.username,u.display_name,COALESCE(u.email,''),u.role,u.avatar_url,
COALESCE(substr(CAST(a.check_in AS TEXT),12,5),''),
COALESCE(substr(CAST(a.check_out AS TEXT),12,5),''),
COALESCE(a.mode,''),
COALESCE((
SELECT r.kind FROM requests r
WHERE r.user_id=u.id AND r.status='approved' AND ? BETWEEN r.start_date AND r.end_date
ORDER BY CASE r.kind WHEN 'leave' THEN 0 ELSE 1 END,r.id DESC
LIMIT 1
),'')
FROM users u
LEFT JOIN attendance a ON a.user_id=u.id AND a.day=?
ORDER BY u.display_name`, day, day)
if err != nil {
return nil, err
}
defer rows.Close()
var roster []DayRosterRow
for rows.Next() {
var row DayRosterRow
if err := rows.Scan(
&row.User.ID, &row.User.Username, &row.User.DisplayName, &row.User.Email,
&row.User.Role, &row.User.AvatarURL, &row.CheckIn, &row.CheckOut,
&row.Mode, &row.RequestKind,
); err != nil {
return nil, err
}
roster = append(roster, row)
}
return roster, rows.Err()
}