-
S
-
+
+
+
+
Step 1 of 4
Welcome
@@ -94,6 +102,8 @@
Welcome
Working… please wait.
+
+
diff --git a/core/engine/templates/totp_login.html b/core/engine/templates/totp_login.html
new file mode 100644
index 00000000..86b23bdc
--- /dev/null
+++ b/core/engine/templates/totp_login.html
@@ -0,0 +1,49 @@
+
+
+
+
+
+
Two-factor authentication — {{.AppName}}
+ {{range .Stylesheets}}
+
+ {{end}}
+
+
+
+
+
+
+
+
+
+
Authentication code
+
6-digit code from your authenticator app.
+ {{if .Error}}
+
{{.Error}}
+ {{end}}
+
+
Back to sign in
+
+
+
+
+
+
diff --git a/core/mail/smtp.go b/core/mail/smtp.go
index eca52a49..e7afedb9 100644
--- a/core/mail/smtp.go
+++ b/core/mail/smtp.go
@@ -4,6 +4,7 @@ package mail
import (
"context"
"crypto/tls"
+ "encoding/base64"
"fmt"
"net"
"net/smtp"
@@ -119,6 +120,60 @@ func sendSMTPS(addr string, auth smtp.Auth, from string, to []string, msg []byte
return client.Quit()
}
+// SendMultipart delivers text/plain and text/html alternative parts.
+func SendMultipart(ctx context.Context, to, subject, textBody, htmlBody string) error {
+ to = strings.TrimSpace(to)
+ if to == "" {
+ return fmt.Errorf("recipient required")
+ }
+ if !Configured() {
+ return fmt.Errorf("smtp not configured")
+ }
+ boundary := "sum-mail-boundary"
+ from := strings.TrimSpace(smtpCfg.From)
+ var b strings.Builder
+ b.WriteString("From: " + from + "\r\n")
+ b.WriteString("To: " + to + "\r\n")
+ b.WriteString("Subject: " + subject + "\r\n")
+ b.WriteString("MIME-Version: 1.0\r\n")
+ b.WriteString("Content-Type: multipart/alternative; boundary=" + boundary + "\r\n\r\n")
+ writePart := func(contentType, body string) {
+ b.WriteString("--" + boundary + "\r\n")
+ b.WriteString("Content-Type: " + contentType + "; charset=UTF-8\r\n")
+ b.WriteString("Content-Transfer-Encoding: base64\r\n\r\n")
+ b.WriteString(base64.StdEncoding.EncodeToString([]byte(body)))
+ b.WriteString("\r\n")
+ }
+ writePart("text/plain", textBody)
+ writePart("text/html", htmlBody)
+ b.WriteString("--" + boundary + "--\r\n")
+ return sendRawMessage(ctx, to, subject, []byte(b.String()))
+}
+
+func sendRawMessage(ctx context.Context, to, subject string, msg []byte) error {
+ addr := fmt.Sprintf("%s:%d", strings.TrimSpace(smtpCfg.Host), smtpCfg.Port)
+ var auth smtp.Auth
+ if u := strings.TrimSpace(smtpCfg.User); u != "" {
+ auth = smtp.PlainAuth("", u, smtpCfg.Password, smtpCfg.Host)
+ }
+ from := strings.TrimSpace(smtpCfg.From)
+ if smtpCfg.Port == 465 {
+ return sendSMTPS(addr, auth, from, []string{to}, msg)
+ }
+ if err := smtp.SendMail(addr, auth, from, []string{to}, msg); err != nil {
+ applog.Warn(ctx, applog.Event{
+ Message: "smtp send failed",
+ Component: "mail",
+ Operation: "send",
+ Status: "failed",
+ Context: map[string]interface{}{"to": to, "subject": subject},
+ Err: err,
+ })
+ return err
+ }
+ return nil
+}
+
// SendPasswordResetEmail notifies a user that an administrator requested a password reset.
func SendPasswordResetEmail(ctx context.Context, to, login, loginURL string) error {
subject := "Sumeru password reset requested"
diff --git a/core/mail/worker.go b/core/mail/worker.go
index bff9df15..28ed0ac5 100644
--- a/core/mail/worker.go
+++ b/core/mail/worker.go
@@ -3,6 +3,7 @@ package mail
import (
"context"
"encoding/json"
+ "strings"
"sumeru/core/applog"
"sumeru/core/queue"
@@ -13,9 +14,10 @@ func init() {
}
type mailJob struct {
- To string `json:"to"`
- Subject string `json:"subject"`
- Body string `json:"body"`
+ To string `json:"to"`
+ Subject string `json:"subject"`
+ Body string `json:"body"`
+ HTMLBody string `json:"html_body,omitempty"`
}
func deliverQueuedMail(ctx context.Context, msg queue.Message) error {
@@ -23,7 +25,13 @@ func deliverQueuedMail(ctx context.Context, msg queue.Message) error {
if err := json.Unmarshal(msg.Payload, &job); err != nil {
return err
}
- if err := Send(ctx, job.To, job.Subject, job.Body); err != nil {
+ var err error
+ if strings.TrimSpace(job.HTMLBody) != "" {
+ err = SendMultipart(ctx, job.To, job.Subject, job.Body, job.HTMLBody)
+ } else {
+ err = Send(ctx, job.To, job.Subject, job.Body)
+ }
+ if err != nil {
applog.Warn(ctx, applog.Event{
Message: "queued mail delivery failed",
Component: "mail",
@@ -40,3 +48,8 @@ func deliverQueuedMail(ctx context.Context, msg queue.Message) error {
func Enqueue(ctx context.Context, to, subject, body string) {
queue.Publish(ctx, "mail", mailJob{To: to, Subject: subject, Body: body})
}
+
+// EnqueueHTML queues multipart HTML + plain alternative email.
+func EnqueueHTML(ctx context.Context, to, subject, htmlBody, textBody string) {
+ queue.Publish(ctx, "mail", mailJob{To: to, Subject: subject, Body: textBody, HTMLBody: htmlBody})
+}
diff --git a/core/modelmeta/model_spec.go b/core/modelmeta/model_spec.go
index 7c462764..35d70c52 100644
--- a/core/modelmeta/model_spec.go
+++ b/core/modelmeta/model_spec.go
@@ -11,6 +11,7 @@ type ModelSpec struct {
Extend bool // true when the struct uses inherit= to extend an existing model
DelegationParent string // set when inherits= names a parent model (_inherits delegation)
CompanyShared bool // company=shared on model tag
+ MailThread bool // mail_thread on model tag
}
// ModelSpecFromStruct reads model= or inherit= from an embedded ModelMeta tag.
@@ -46,9 +47,19 @@ func ModelSpecFromTags(tags FieldTags, goName string) (ModelSpec, error) {
return ModelSpec{Name: "-", Extend: false}, nil
}
if tags.Model != "" {
- return ModelSpec{Name: tags.Model, Extend: false, CompanyShared: tags.Company == "shared"}, nil
+ return ModelSpec{
+ Name: tags.Model,
+ Extend: false,
+ CompanyShared: tags.Company == "shared",
+ MailThread: tags.MailThread,
+ }, nil
}
- return ModelSpec{Name: ModelNameFromGo(goName), Extend: false, CompanyShared: tags.Company == "shared"}, nil
+ return ModelSpec{
+ Name: ModelNameFromGo(goName),
+ Extend: false,
+ CompanyShared: tags.Company == "shared",
+ MailThread: tags.MailThread,
+ }, nil
}
// ModelNameFromStruct reads the technical model name from an embedded ModelMeta tag,
diff --git a/core/modelmeta/tags.go b/core/modelmeta/tags.go
index 5f57659c..d56cccdb 100644
--- a/core/modelmeta/tags.go
+++ b/core/modelmeta/tags.go
@@ -44,6 +44,7 @@ type FieldTags struct {
Related string
Compute string
Company string // model tag: company=shared skips auto multi-company isolation
+ MailThread bool // model tag: mail_thread — chatter auto-subscribe hooks
}
// ParseModelTag parses the sumeru tag on an embedded ModelMeta.
@@ -241,6 +242,8 @@ func setTagOption(tags *FieldTags, key, value string) error {
tags.Related = value
case "compute":
tags.Compute = value
+ case "mail_thread":
+ tags.MailThread = true
default:
return fmt.Errorf("unknown sumeru tag %q", key)
}
diff --git a/core/modelreg/activate.go b/core/modelreg/activate.go
index 8281b89f..6cd8a70f 100644
--- a/core/modelreg/activate.go
+++ b/core/modelreg/activate.go
@@ -11,6 +11,7 @@ type pendingModel struct {
name string
extend bool
companyShared bool
+ mailThread bool
fields []orm.FieldDefinition
}
@@ -64,6 +65,9 @@ func activateAllLocked(moduleOrder []string) error {
if pending.companyShared {
orm.SetModelCompanyShared(pending.name, true)
}
+ if pending.mailThread {
+ orm.SetModelMailThread(pending.name, true)
+ }
}
}
return nil
diff --git a/core/modelreg/register.go b/core/modelreg/register.go
index 72cf34e4..abfb3187 100644
--- a/core/modelreg/register.go
+++ b/core/modelreg/register.go
@@ -48,6 +48,7 @@ func MustRegister(module string, models ...any) {
name: entry.spec.Name,
extend: entry.spec.Extend,
companyShared: entry.spec.CompanyShared,
+ mailThread: entry.spec.MailThread,
fields: fields,
})
}
diff --git a/core/orm/bus_notify.go b/core/orm/bus_notify.go
new file mode 100644
index 00000000..6543921d
--- /dev/null
+++ b/core/orm/bus_notify.go
@@ -0,0 +1,18 @@
+package orm
+
+import (
+ "fmt"
+ "strconv"
+)
+
+// NotifyBusEvent sends PostgreSQL NOTIFY for multi-process fan-out.
+func NotifyBusEvent(eventID int64) error {
+ if DB == nil || eventID <= 0 {
+ return nil
+ }
+ _, err := DB.Exec(`SELECT pg_notify($1, $2)`, BusNotifyChannel(), strconv.FormatInt(eventID, 10))
+ if err != nil {
+ return fmt.Errorf("pg_notify bus: %w", err)
+ }
+ return nil
+}
diff --git a/core/orm/bus_publish_hook.go b/core/orm/bus_publish_hook.go
new file mode 100644
index 00000000..f3501113
--- /dev/null
+++ b/core/orm/bus_publish_hook.go
@@ -0,0 +1,28 @@
+package orm
+
+import (
+ "context"
+ "strconv"
+)
+
+var busPublishHook func(ctx context.Context, channel string, payload map[string]interface{})
+
+// SetBusPublishHook registers the server implementation (WebSocket fan-out).
+func SetBusPublishHook(fn func(ctx context.Context, channel string, payload map[string]interface{})) {
+ busPublishHook = fn
+}
+
+// PublishBusChannel emits a bus event when a hook is registered.
+func PublishBusChannel(ctx context.Context, channel string, payload map[string]interface{}) {
+ if busPublishHook != nil && channel != "" {
+ busPublishHook(ctx, channel, payload)
+ }
+}
+
+// UserNotificationsBusChannel returns the inbox channel for uid.
+func UserNotificationsBusChannel(uid int) string {
+ if uid <= 0 {
+ return ""
+ }
+ return "user/" + strconv.Itoa(uid) + "/notifications"
+}
diff --git a/core/orm/mail_thread.go b/core/orm/mail_thread.go
new file mode 100644
index 00000000..1cf5fc40
--- /dev/null
+++ b/core/orm/mail_thread.go
@@ -0,0 +1,26 @@
+package orm
+
+import "sync"
+
+var (
+ mailThreadMu sync.RWMutex
+ mailThreadModels = map[string]bool{}
+)
+
+// SetModelMailThread marks a model as chatter-capable (auto-subscribe on create).
+func SetModelMailThread(modelName string, enabled bool) {
+ mailThreadMu.Lock()
+ defer mailThreadMu.Unlock()
+ if enabled {
+ mailThreadModels[modelName] = true
+ return
+ }
+ delete(mailThreadModels, modelName)
+}
+
+// ModelHasMailThread reports whether model uses mail thread hooks.
+func ModelHasMailThread(modelName string) bool {
+ mailThreadMu.RLock()
+ defer mailThreadMu.RUnlock()
+ return mailThreadModels[modelName]
+}
diff --git a/core/orm/sys_bus_event.go b/core/orm/sys_bus_event.go
new file mode 100644
index 00000000..b7c613f9
--- /dev/null
+++ b/core/orm/sys_bus_event.go
@@ -0,0 +1,136 @@
+package orm
+
+import (
+ "context"
+ "encoding/json"
+ "fmt"
+ "strings"
+ "time"
+
+ "sumeru/core/modelmeta"
+)
+
+type SysBusEvent struct {
+ modelmeta.ModelMeta `sumeru:"model=sys.bus.event"`
+
+ Channel modelmeta.String `sumeru:"required,index"`
+ PayloadJson modelmeta.Text `sumeru:"column=payload_json"`
+ CreatedAt modelmeta.DateTime `sumeru:"required,index,column=created_at"`
+}
+
+const busNotifyChannel = "sumeru_bus"
+
+// BusNotifyChannel returns the PostgreSQL NOTIFY channel name for bus fan-out.
+func BusNotifyChannel() string {
+ return busNotifyChannel
+}
+
+// InsertBusEvent persists an event and returns its id (0 if DB unavailable).
+func InsertBusEvent(ctx context.Context, channel string, payload map[string]interface{}) (int64, error) {
+ if DB == nil || channel == "" {
+ return 0, fmt.Errorf("bus event: database or channel missing")
+ }
+ if _, ok := Registry["sys.bus.event"]; !ok {
+ return 0, fmt.Errorf("sys.bus.event model not registered")
+ }
+ pj := ""
+ if payload != nil {
+ b, err := json.Marshal(payload)
+ if err != nil {
+ return 0, err
+ }
+ pj = string(b)
+ }
+ inst := Registry["sys.bus.event"]
+ id, err := Create(ctx, inst, map[string]interface{}{
+ "channel": channel,
+ "payload_json": pj,
+ "created_at": time.Now().UTC().Format(time.RFC3339),
+ })
+ if err != nil {
+ return 0, err
+ }
+ return int64(id), nil
+}
+
+// LoadBusEventByID loads channel and payload for a persisted event id.
+func LoadBusEventByID(ctx context.Context, eventID int64) (channel string, payload map[string]interface{}, err error) {
+ if DB == nil || eventID <= 0 {
+ return "", nil, fmt.Errorf("invalid bus event load")
+ }
+ tbl, ok := busEventTable()
+ if !ok {
+ return "", nil, fmt.Errorf("sys.bus.event table missing")
+ }
+ var ch string
+ var pj string
+ err = DB.QueryRowContext(ctx, `SELECT channel, COALESCE(payload_json,'') FROM `+tbl+` WHERE id = $1`, eventID).Scan(&ch, &pj)
+ if err != nil {
+ return "", nil, err
+ }
+ payload = map[string]interface{}{}
+ if pj != "" {
+ _ = json.Unmarshal([]byte(pj), &payload)
+ }
+ return ch, payload, nil
+}
+
+// ListBusEventsAfter returns events for channels with id > afterID (cap limit).
+func ListBusEventsAfter(ctx context.Context, channels []string, afterID int64, limit int) ([]BusEventRow, error) {
+ if DB == nil || len(channels) == 0 {
+ return nil, nil
+ }
+ tbl, ok := busEventTable()
+ if !ok {
+ return nil, nil
+ }
+ if limit <= 0 || limit > 500 {
+ limit = 500
+ }
+ // Build IN clause for channels — ponytail: small channel list per subscribe frame.
+ placeholders := make([]string, len(channels))
+ args := make([]interface{}, 0, len(channels)+1)
+ args = append(args, afterID)
+ for i, ch := range channels {
+ placeholders[i] = fmt.Sprintf("$%d", i+2)
+ args = append(args, ch)
+ }
+ inClause := strings.Join(placeholders, ",")
+ query := fmt.Sprintf(
+ `SELECT id, channel, COALESCE(payload_json,'') FROM `+tbl+` WHERE id > $1 AND channel IN (%s) ORDER BY id ASC LIMIT %d`,
+ inClause,
+ limit,
+ )
+ rows, err := DB.QueryContext(ctx, query, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ var out []BusEventRow
+ for rows.Next() {
+ var row BusEventRow
+ var pj string
+ if err := rows.Scan(&row.ID, &row.Channel, &pj); err != nil {
+ return out, err
+ }
+ row.Payload = map[string]interface{}{}
+ if pj != "" {
+ _ = json.Unmarshal([]byte(pj), &row.Payload)
+ }
+ out = append(out, row)
+ }
+ return out, rows.Err()
+}
+
+type BusEventRow struct {
+ ID int64
+ Channel string
+ Payload map[string]interface{}
+}
+
+func busEventTable() (string, bool) {
+ if _, ok := Registry["sys.bus.event"]; !ok {
+ return "", false
+ }
+ return MustQuotedTableName("sys.bus.event"), true
+}
diff --git a/core/orm/user_totp.go b/core/orm/user_totp.go
new file mode 100644
index 00000000..732b8624
--- /dev/null
+++ b/core/orm/user_totp.go
@@ -0,0 +1,146 @@
+package orm
+
+import (
+ "context"
+ "database/sql"
+ "fmt"
+ "strings"
+
+ "sumeru/core/server/auth"
+)
+
+// BeginOwnTOTPEnrollment stores a new secret with totp_enabled=false for the given user.
+func BeginOwnTOTPEnrollment(ctx context.Context, actor, userID int) (secret string, err error) {
+ if err := assertOwnUser(actor, userID); err != nil {
+ return "", err
+ }
+ secret, err = auth.GenerateTOTPSecret()
+ if err != nil {
+ return "", err
+ }
+ return secret, writeUserTOTPSecret(ctx, userID, secret, false)
+}
+
+// ConfirmOwnTOTPEnrollment validates code against stored secret and enables TOTP.
+func ConfirmOwnTOTPEnrollment(ctx context.Context, actor, userID int, code string) error {
+ if err := assertOwnUser(actor, userID); err != nil {
+ return err
+ }
+ secret, enabled, err := readUserTOTP(ctx, userID)
+ if err != nil {
+ return err
+ }
+ if enabled {
+ return fmt.Errorf("two-factor authentication is already enabled")
+ }
+ if secret == "" {
+ return fmt.Errorf("start enrollment before confirming")
+ }
+ if !auth.ValidateTOTP(secret, code) {
+ return fmt.Errorf("invalid authentication code")
+ }
+ return writeUserTOTPSecret(ctx, userID, secret, true)
+}
+
+// DisableOwnTOTP requires a valid code and clears TOTP for the user.
+func DisableOwnTOTP(ctx context.Context, actor, userID int, code string) error {
+ if err := assertOwnUser(actor, userID); err != nil {
+ return err
+ }
+ secret, enabled, err := readUserTOTP(ctx, userID)
+ if err != nil {
+ return err
+ }
+ if !enabled || secret == "" {
+ return fmt.Errorf("two-factor authentication is not enabled")
+ }
+ if !auth.ValidateTOTP(secret, code) {
+ return fmt.Errorf("invalid authentication code")
+ }
+ return writeUserTOTPSecret(ctx, userID, "", false)
+}
+
+// DisableUserTOTP clears TOTP for userID (system administrator).
+func DisableUserTOTP(ctx context.Context, actor, userID int) error {
+ if userID <= 0 {
+ return fmt.Errorf("invalid user id")
+ }
+ if actor <= 0 {
+ return fmt.Errorf("unauthenticated")
+ }
+ if !UserHasGroupXML(ctx, actor, "base.group_system") {
+ return fmt.Errorf("two-factor reset requires system administrator")
+ }
+ return writeUserTOTPSecret(ctx, userID, "", false)
+}
+
+// UserTOTPEnabled reports whether TOTP login is required for userID.
+func UserTOTPEnabled(ctx context.Context, userID int) (bool, error) {
+ _, enabled, err := readUserTOTP(ctx, userID)
+ return enabled, err
+}
+
+// UserTOTPCredentialsForLogin returns secret/enabled for MFA verification (kernel read).
+func UserTOTPCredentialsForLogin(ctx context.Context, userID int) (secret string, enabled bool) {
+ secret, enabled, err := readUserTOTP(ctx, userID)
+ if err != nil {
+ return "", false
+ }
+ return secret, enabled
+}
+
+// OwnTOTPEnrollmentPending returns the pending secret when enrollment started but not confirmed.
+func OwnTOTPEnrollmentPending(ctx context.Context, actor, userID int) (secret string, ok bool, err error) {
+ if err := assertOwnUser(actor, userID); err != nil {
+ return "", false, err
+ }
+ secret, enabled, err := readUserTOTP(ctx, userID)
+ if err != nil {
+ return "", false, err
+ }
+ if enabled || secret == "" {
+ return "", false, nil
+ }
+ return secret, true, nil
+}
+
+func assertOwnUser(actor, userID int) error {
+ if userID <= 0 || actor <= 0 {
+ return fmt.Errorf("unauthenticated")
+ }
+ if actor != userID {
+ return fmt.Errorf("two-factor enrollment requires your own account")
+ }
+ return nil
+}
+
+func readUserTOTP(ctx context.Context, userID int) (secret string, enabled bool, err error) {
+ if DB == nil {
+ return "", false, fmt.Errorf("database unavailable")
+ }
+ bypass := AuditedBypass(ctx, "user.totp.read")
+ tbl := MustQuotedTableName("core.user")
+ var secretNull sql.NullString
+ var enabledVal bool
+ err = DB.QueryRowContext(bypass,
+ `SELECT COALESCE(totp_secret, ''), COALESCE(totp_enabled, false) FROM `+tbl+` WHERE id = $1`,
+ userID,
+ ).Scan(&secretNull, &enabledVal)
+ if err == sql.ErrNoRows {
+ return "", false, fmt.Errorf("user not found")
+ }
+ if err != nil {
+ return "", false, err
+ }
+ return strings.TrimSpace(secretNull.String), enabledVal, nil
+}
+
+func writeUserTOTPSecret(ctx context.Context, userID int, secret string, enabled bool) error {
+ values := map[string]interface{}{
+ "totp_secret": strings.TrimSpace(secret),
+ "totp_enabled": enabled,
+ }
+ return WithElevated(ctx, "user.totp.write", func(bypass context.Context) error {
+ return UpdateRecordByID(bypass, "core.user", userID, values)
+ })
+}
diff --git a/core/ormmodels/zmodels.go b/core/ormmodels/zmodels.go
index 339b83a6..5290d57f 100644
--- a/core/ormmodels/zmodels.go
+++ b/core/ormmodels/zmodels.go
@@ -14,6 +14,7 @@ func init() {
&orm.SysActionURL{},
&orm.SysActionWindow{},
&orm.SysApprovalRule{},
+ &orm.SysBusEvent{},
&orm.SysField{},
&orm.SysMenu{},
&orm.SysModel{},
diff --git a/core/security/fields.go b/core/security/fields.go
index 928db559..17763f49 100644
--- a/core/security/fields.go
+++ b/core/security/fields.go
@@ -24,6 +24,12 @@ var FieldRegistry = map[string]map[string]FieldPolicy{
"sys.attachment": {
"datas": {ReadRedact: true},
},
+ "sys.auth.provider": {
+ "client_secret": {ReadRedact: true, WriteDenyUnlessSys: true},
+ },
+ "core.user.trusteddevice": {
+ "token_hash": {ReadRedact: true, WriteDenyUnlessSys: true},
+ },
}
// ReadRedactFields returns fields stripped on read for a model.
diff --git a/core/server/auth/oidc_verify.go b/core/server/auth/oidc_verify.go
new file mode 100644
index 00000000..675246d3
--- /dev/null
+++ b/core/server/auth/oidc_verify.go
@@ -0,0 +1,320 @@
+package auth
+
+import (
+ "context"
+ "crypto"
+ "crypto/ecdsa"
+ "crypto/elliptic"
+ "crypto/rsa"
+ "encoding/base64"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "math/big"
+ "net/http"
+ "strings"
+ "sync"
+ "time"
+)
+
+type jwksCacheEntry struct {
+ keys map[string]crypto.PublicKey
+ expires time.Time
+}
+
+var (
+ jwksCacheMu sync.Mutex
+ jwksCache = map[string]jwksCacheEntry{}
+ jwksTTL = 15 * time.Minute
+)
+
+// VerifyIDToken validates an OIDC id_token signature and standard claims when jwksURL is set.
+func VerifyIDToken(ctx context.Context, idToken, issuer, clientID, jwksURL string) (map[string]interface{}, error) {
+ idToken = strings.TrimSpace(idToken)
+ if idToken == "" {
+ return nil, errors.New("id_token required")
+ }
+ parts := strings.Split(idToken, ".")
+ if len(parts) != 3 {
+ return nil, errors.New("malformed id_token")
+ }
+ headerJSON, err := base64.RawURLEncoding.DecodeString(parts[0])
+ if err != nil {
+ return nil, fmt.Errorf("id_token header: %w", err)
+ }
+ var header struct {
+ Alg string `json:"alg"`
+ Kid string `json:"kid"`
+ }
+ if err := json.Unmarshal(headerJSON, &header); err != nil {
+ return nil, fmt.Errorf("id_token header json: %w", err)
+ }
+ alg := strings.ToUpper(strings.TrimSpace(header.Alg))
+ if alg != "RS256" && alg != "ES256" {
+ return nil, fmt.Errorf("unsupported jwt alg %q", header.Alg)
+ }
+ jwksURL = strings.TrimSpace(jwksURL)
+ if jwksURL == "" {
+ return nil, errors.New("jwks_url required for verification")
+ }
+ pub, err := publicKeyForJWKS(ctx, jwksURL, header.Kid, alg)
+ if err != nil {
+ return nil, err
+ }
+ signingInput := parts[0] + "." + parts[1]
+ sig, err := base64.RawURLEncoding.DecodeString(parts[2])
+ if err != nil {
+ return nil, fmt.Errorf("id_token signature: %w", err)
+ }
+ hash := crypto.SHA256.New()
+ _, _ = hash.Write([]byte(signingInput))
+ digest := hash.Sum(nil)
+ switch key := pub.(type) {
+ case *rsa.PublicKey:
+ if alg != "RS256" {
+ return nil, fmt.Errorf("alg/key mismatch")
+ }
+ if err := rsa.VerifyPKCS1v15(key, crypto.SHA256, digest, sig); err != nil {
+ return nil, fmt.Errorf("rsa verify: %w", err)
+ }
+ case *ecdsa.PublicKey:
+ if alg != "ES256" {
+ return nil, fmt.Errorf("alg/key mismatch")
+ }
+ if !ecdsa.VerifyASN1(key, digest, sig) {
+ return nil, errors.New("ecdsa verify failed")
+ }
+ default:
+ return nil, errors.New("unsupported public key type")
+ }
+ payload, err := base64.RawURLEncoding.DecodeString(parts[1])
+ if err != nil {
+ return nil, fmt.Errorf("id_token payload: %w", err)
+ }
+ var claims map[string]interface{}
+ if err := json.Unmarshal(payload, &claims); err != nil {
+ return nil, fmt.Errorf("id_token claims: %w", err)
+ }
+ if err := validateIDTokenClaims(claims, issuer, clientID); err != nil {
+ return nil, err
+ }
+ return claims, nil
+}
+
+func validateIDTokenClaims(claims map[string]interface{}, issuer, clientID string) error {
+ now := time.Now().Unix()
+ if exp, ok := claimInt64(claims["exp"]); ok && exp > 0 && now > exp+60 {
+ return errors.New("id_token expired")
+ }
+ issuer = strings.TrimSpace(issuer)
+ if issuer != "" {
+ if iss, _ := claims["iss"].(string); strings.TrimSpace(iss) != issuer {
+ return fmt.Errorf("issuer mismatch")
+ }
+ }
+ clientID = strings.TrimSpace(clientID)
+ if clientID != "" {
+ if !audienceContains(claims["aud"], clientID) {
+ return fmt.Errorf("audience mismatch")
+ }
+ }
+ if sub, _ := claims["sub"].(string); strings.TrimSpace(sub) == "" {
+ return errors.New("missing sub")
+ }
+ return nil
+}
+
+func audienceContains(aud interface{}, clientID string) bool {
+ switch v := aud.(type) {
+ case string:
+ return strings.TrimSpace(v) == clientID
+ case []interface{}:
+ for _, item := range v {
+ if s, ok := item.(string); ok && strings.TrimSpace(s) == clientID {
+ return true
+ }
+ }
+ }
+ return false
+}
+
+func claimInt64(v interface{}) (int64, bool) {
+ switch n := v.(type) {
+ case float64:
+ return int64(n), true
+ case int64:
+ return n, true
+ case int:
+ return int64(n), true
+ default:
+ return 0, false
+ }
+}
+
+func publicKeyForJWKS(ctx context.Context, jwksURL, kid, alg string) (crypto.PublicKey, error) {
+ jwksCacheMu.Lock()
+ if ent, ok := jwksCache[jwksURL]; ok && time.Now().Before(ent.expires) {
+ if kid != "" {
+ if k, ok := ent.keys[kid]; ok {
+ jwksCacheMu.Unlock()
+ return k, nil
+ }
+ } else if len(ent.keys) == 1 {
+ for _, k := range ent.keys {
+ jwksCacheMu.Unlock()
+ return k, nil
+ }
+ }
+ }
+ jwksCacheMu.Unlock()
+
+ keys, err := fetchJWKS(ctx, jwksURL)
+ if err != nil {
+ return nil, err
+ }
+ jwksCacheMu.Lock()
+ jwksCache[jwksURL] = jwksCacheEntry{keys: keys, expires: time.Now().Add(jwksTTL)}
+ jwksCacheMu.Unlock()
+
+ if kid != "" {
+ if k, ok := keys[kid]; ok {
+ return k, nil
+ }
+ return nil, fmt.Errorf("jwks kid %q not found", kid)
+ }
+ if len(keys) == 1 {
+ for _, k := range keys {
+ return k, nil
+ }
+ }
+ return nil, errors.New("jwks kid required")
+}
+
+func fetchJWKS(ctx context.Context, jwksURL string) (map[string]crypto.PublicKey, error) {
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, jwksURL, nil)
+ if err != nil {
+ return nil, err
+ }
+ client := &http.Client{Timeout: 10 * time.Second}
+ resp, err := client.Do(req)
+ if err != nil {
+ return nil, err
+ }
+ defer resp.Body.Close()
+ body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
+ if err != nil {
+ return nil, err
+ }
+ if resp.StatusCode >= 400 {
+ return nil, fmt.Errorf("jwks http %d", resp.StatusCode)
+ }
+ var doc struct {
+ Keys []json.RawMessage `json:"keys"`
+ }
+ if err := json.Unmarshal(body, &doc); err != nil {
+ return nil, err
+ }
+ out := make(map[string]crypto.PublicKey, len(doc.Keys))
+ for _, raw := range doc.Keys {
+ var meta struct {
+ Kty string `json:"kty"`
+ Kid string `json:"kid"`
+ Alg string `json:"alg"`
+ Use string `json:"use"`
+ N string `json:"n"`
+ E string `json:"e"`
+ Crv string `json:"crv"`
+ X string `json:"x"`
+ Y string `json:"y"`
+ }
+ if err := json.Unmarshal(raw, &meta); err != nil {
+ continue
+ }
+ if meta.Use != "" && meta.Use != "sig" {
+ continue
+ }
+ kid := meta.Kid
+ if kid == "" {
+ kid = "_"
+ }
+ switch meta.Kty {
+ case "RSA":
+ pub, err := rsaPublicFromJWK(meta.N, meta.E)
+ if err != nil {
+ continue
+ }
+ out[kid] = pub
+ case "EC":
+ pub, err := ecPublicFromJWK(meta.Crv, meta.X, meta.Y)
+ if err != nil {
+ continue
+ }
+ out[kid] = pub
+ }
+ }
+ if len(out) == 0 {
+ return nil, errors.New("no usable jwks keys")
+ }
+ return out, nil
+}
+
+func rsaPublicFromJWK(nB64, eB64 string) (*rsa.PublicKey, error) {
+ nBytes, err := base64.RawURLEncoding.DecodeString(nB64)
+ if err != nil {
+ return nil, err
+ }
+ eBytes, err := base64.RawURLEncoding.DecodeString(eB64)
+ if err != nil {
+ return nil, err
+ }
+ var eInt int
+ for _, b := range eBytes {
+ eInt = eInt<<8 + int(b)
+ }
+ if eInt == 0 {
+ eInt = 65537
+ }
+ return &rsa.PublicKey{N: new(big.Int).SetBytes(nBytes), E: eInt}, nil
+}
+
+func ecPublicFromJWK(crv, xB64, yB64 string) (*ecdsa.PublicKey, error) {
+ if crv != "P-256" {
+ return nil, fmt.Errorf("unsupported crv %q", crv)
+ }
+ xBytes, err := base64.RawURLEncoding.DecodeString(xB64)
+ if err != nil {
+ return nil, err
+ }
+ yBytes, err := base64.RawURLEncoding.DecodeString(yB64)
+ if err != nil {
+ return nil, err
+ }
+ // ponytail: JWK EC coordinates via big.Int until stdlib exposes ParseUncompressed for OIDC JWKS.
+ pub := &ecdsa.PublicKey{Curve: elliptic.P256()} //nolint:staticcheck // SA1019 coordinate fill for JWKS
+ pub.X = new(big.Int).SetBytes(xBytes) //nolint:staticcheck // SA1019
+ pub.Y = new(big.Int).SetBytes(yBytes) //nolint:staticcheck // SA1019
+ return pub, nil
+}
+
+// ParseIDTokenClaimsUnsafe decodes id_token payload without signature verification (dev or post-verify).
+func ParseIDTokenClaimsUnsafe(idToken string) (email, sub string, claims map[string]interface{}) {
+ s := strings.TrimSpace(idToken)
+ if s == "" || !strings.Contains(s, ".") {
+ return "", "", nil
+ }
+ parts := strings.Split(s, ".")
+ if len(parts) < 2 {
+ return "", "", nil
+ }
+ payload, err := base64.RawURLEncoding.DecodeString(parts[1])
+ if err != nil {
+ return "", "", nil
+ }
+ if json.Unmarshal(payload, &claims) != nil {
+ return "", "", nil
+ }
+ email, _ = claims["email"].(string)
+ sub, _ = claims["sub"].(string)
+ return email, sub, claims
+}
diff --git a/core/server/auth/totp.go b/core/server/auth/totp.go
new file mode 100644
index 00000000..f4101e0e
--- /dev/null
+++ b/core/server/auth/totp.go
@@ -0,0 +1,52 @@
+package auth
+
+import (
+ "crypto/hmac"
+ "crypto/rand"
+ "crypto/sha1"
+ "encoding/base32"
+ "encoding/binary"
+ "fmt"
+ "strings"
+ "time"
+)
+
+// GenerateTOTPSecret returns a base32 secret suitable for authenticator apps.
+func GenerateTOTPSecret() (string, error) {
+ buf := make([]byte, 20)
+ if _, err := rand.Read(buf); err != nil {
+ return "", err
+ }
+ return base32.StdEncoding.WithPadding(base32.NoPadding).EncodeToString(buf), nil
+}
+
+// ValidateTOTP checks code against secret with ±1 step window.
+func ValidateTOTP(secret, code string) bool {
+ secret = strings.TrimSpace(strings.ToUpper(secret))
+ code = strings.TrimSpace(code)
+ if secret == "" || len(code) != 6 {
+ return false
+ }
+ key, err := base32.StdEncoding.WithPadding(base32.NoPadding).DecodeString(secret)
+ if err != nil {
+ return false
+ }
+ now := time.Now().UTC().Unix() / 30
+ for _, step := range []int64{now - 1, now, now + 1} {
+ if fmt.Sprintf("%06d", hotp(key, step)) == code {
+ return true
+ }
+ }
+ return false
+}
+
+func hotp(key []byte, counter int64) int32 {
+ var buf [8]byte
+ binary.BigEndian.PutUint64(buf[:], uint64(counter))
+ mac := hmac.New(sha1.New, key)
+ _, _ = mac.Write(buf[:])
+ sum := mac.Sum(nil)
+ offset := sum[len(sum)-1] & 0x0f
+ truncated := binary.BigEndian.Uint32(sum[offset:offset+4]) & 0x7fffffff
+ return int32(truncated % 1000000)
+}
diff --git a/core/server/auth/totp_qr.go b/core/server/auth/totp_qr.go
new file mode 100644
index 00000000..bb7ac8dd
--- /dev/null
+++ b/core/server/auth/totp_qr.go
@@ -0,0 +1,24 @@
+package auth
+
+import (
+ "fmt"
+ "net/url"
+ "strings"
+)
+
+// TOTPOtpauthURI builds a standard otpauth URI for authenticator apps.
+func TOTPOtpauthURI(issuer, accountName, secret string) string {
+ issuer = strings.TrimSpace(issuer)
+ accountName = strings.TrimSpace(accountName)
+ secret = strings.TrimSpace(secret)
+ label := url.PathEscape(issuer + ":" + accountName)
+ if issuer == "" {
+ label = url.PathEscape(accountName)
+ }
+ q := url.Values{}
+ q.Set("secret", secret)
+ if issuer != "" {
+ q.Set("issuer", issuer)
+ }
+ return fmt.Sprintf("otpauth://totp/%s?%s", label, q.Encode())
+}
diff --git a/core/server/auth/totp_qr_png.go b/core/server/auth/totp_qr_png.go
new file mode 100644
index 00000000..bd564a1a
--- /dev/null
+++ b/core/server/auth/totp_qr_png.go
@@ -0,0 +1,11 @@
+package auth
+
+import qrcode "github.com/skip2/go-qrcode"
+
+// TOTPQRPng returns a PNG image for an otpauth URI.
+func TOTPQRPng(otpauthURI string, size int) ([]byte, error) {
+ if size <= 0 {
+ size = 200
+ }
+ return qrcode.Encode(otpauthURI, qrcode.Medium, size)
+}
diff --git a/core/server/run.go b/core/server/run.go
index 8cc0185b..5cb0d234 100644
--- a/core/server/run.go
+++ b/core/server/run.go
@@ -196,6 +196,7 @@ func Run() {
defer stop()
scheduler.Start(rootCtx, time.Minute)
orm.StartOutboxDrain(rootCtx, 5*time.Second)
+ web.StartBusListener(rootCtx, databaseSource)
listenHost := listenAddr(config.AppConfig.HttpInterface, config.AppConfig.HttpPort)
applog.InfoMsg(ctx, "server", "listen", "Server starting",
diff --git a/core/server/web/apikey_flash.go b/core/server/web/apikey_flash.go
deleted file mode 100644
index 4335fe41..00000000
--- a/core/server/web/apikey_flash.go
+++ /dev/null
@@ -1,22 +0,0 @@
-package web
-
-import (
- "net/http"
-)
-
-const apiKeyFlashCookie = "sumeru_api_key_flash"
-
-// SetAPIKeyFlash stores a one-time raw API key in an HttpOnly cookie (never in redirect URLs).
-func SetAPIKeyFlash(w http.ResponseWriter, raw string) {
- setNamedCookie(w, apiKeyFlashCookie, raw, "/", 120, http.SameSiteLaxMode)
-}
-
-// ConsumeAPIKeyFlash reads and clears the one-time API key flash cookie.
-func ConsumeAPIKeyFlash(r *http.Request, w http.ResponseWriter) string {
- c, err := r.Cookie(apiKeyFlashCookie)
- clearNamedCookie(w, apiKeyFlashCookie, "/", http.SameSiteLaxMode)
- if err != nil || c.Value == "" {
- return ""
- }
- return c.Value
-}
diff --git a/core/server/web/apikey_create.go b/core/server/web/apikey_handlers.go
similarity index 53%
rename from core/server/web/apikey_create.go
rename to core/server/web/apikey_handlers.go
index 104b1b62..f5766cbb 100644
--- a/core/server/web/apikey_create.go
+++ b/core/server/web/apikey_handlers.go
@@ -11,11 +11,13 @@ import (
const apiKeyModel = "core.user.apikey"
+const apiKeyRevealRoute = "/web/apikey/reveal"
+
+func registerAPIKeyRevealRoute() {
+ registerSession(http.MethodGet, apiKeyRevealRoute, APIKeyRevealHandler)
+}
+
// ActionCreateAPIKey generates a one-time raw API key for a user.
-//
-// Form fields: name (optional label), user_id (optional; defaults to signed-in user), next (redirect target).
-// The raw key is never placed in the redirect URL — it is delivered once via SetAPIKeyFlash
-// and shown on the next page by ConsumePageFlashes (see page_flash.go).
func ActionCreateAPIKey(w http.ResponseWriter, r *http.Request) {
if !requireLoginAndPOST(w, r) {
return
@@ -31,7 +33,7 @@ func ActionCreateAPIKey(w http.ResponseWriter, r *http.Request) {
rawKey, err := orm.CreateAPIKeyForUser(ctx, targetUserID, keyName)
if err != nil {
WebLogEvent(ctx, WebLogInput{
- Route: "/web/action/create_api_key",
+ Route: createAPIKeyRoute,
Message: "Could not create API key",
Code: errcode.InternalError,
Operation: "create_api_key",
@@ -49,8 +51,6 @@ func ActionCreateAPIKey(w http.ResponseWriter, r *http.Request) {
redirectWithWebMessage(w, r, r.PostFormValue("next"), "api_key_created")
}
-// apiKeyTargetUserID resolves which user receives the new key; falls back to the session user.
-// Targeting another user requires base.group_system.
func apiKeyTargetUserID(r *http.Request) int {
sessionUID := SessionUserID(r)
userID, _ := strconv.Atoi(strings.TrimSpace(r.PostFormValue("user_id")))
@@ -62,3 +62,27 @@ func apiKeyTargetUserID(r *http.Request) int {
}
return sessionUID
}
+
+// APIKeyRevealHandler shows a one-time API key from the flash cookie (never embedded in shell HTML).
+func APIKeyRevealHandler(w http.ResponseWriter, r *http.Request) {
+ if !requireLogin(w, r) {
+ return
+ }
+ raw := ConsumeAPIKeyFlash(r, w)
+ if raw == "" {
+ http.Error(w, "No API key pending display", http.StatusNotFound)
+ return
+ }
+ w.Header().Set("Content-Type", "text/html; charset=utf-8")
+ w.Header().Set("X-Content-Type-Options", "nosniff")
+ _, _ = w.Write([]byte(`
API Key`))
+ _, _ = w.Write([]byte(`
API key (shown once)
Copy this key now. It will not be shown again.
`))
+ _, _ = w.Write([]byte(`
`))
+ _, _ = w.Write([]byte(htmlEscape(raw)))
+ _, _ = w.Write([]byte(`
Back to app
`))
+}
+
+func htmlEscape(s string) string {
+ repl := strings.NewReplacer("&", "&", "<", "<", ">", ">", `"`, """)
+ return repl.Replace(s)
+}
diff --git a/core/server/web/apikey_reveal.go b/core/server/web/apikey_reveal.go
deleted file mode 100644
index 39e7274a..00000000
--- a/core/server/web/apikey_reveal.go
+++ /dev/null
@@ -1,36 +0,0 @@
-package web
-
-import (
- "net/http"
- "strings"
-)
-
-const apiKeyRevealRoute = "/web/apikey/reveal"
-
-func registerAPIKeyRevealRoute() {
- registerSession(http.MethodGet, apiKeyRevealRoute, APIKeyRevealHandler)
-}
-
-// APIKeyRevealHandler shows a one-time API key from the flash cookie (never embedded in shell HTML).
-func APIKeyRevealHandler(w http.ResponseWriter, r *http.Request) {
- if !requireLogin(w, r) {
- return
- }
- raw := ConsumeAPIKeyFlash(r, w)
- if raw == "" {
- http.Error(w, "No API key pending display", http.StatusNotFound)
- return
- }
- w.Header().Set("Content-Type", "text/html; charset=utf-8")
- w.Header().Set("X-Content-Type-Options", "nosniff")
- _, _ = w.Write([]byte(`
API Key`))
- _, _ = w.Write([]byte(`
API key (shown once)
Copy this key now. It will not be shown again.
`))
- _, _ = w.Write([]byte(`
`))
- _, _ = w.Write([]byte(htmlEscape(raw)))
- _, _ = w.Write([]byte(`
Back to app
`))
-}
-
-func htmlEscape(s string) string {
- repl := strings.NewReplacer("&", "&", "<", "<", ">", ">", `"`, """)
- return repl.Replace(s)
-}
diff --git a/core/server/web/apps_flash.go b/core/server/web/apps_flash.go
index 88b3df64..952b7d3f 100644
--- a/core/server/web/apps_flash.go
+++ b/core/server/web/apps_flash.go
@@ -72,7 +72,7 @@ func appsFlashFromMessage(msg string, displayNames map[string]string) (render.Fl
return render.FlashMessage{Kind: "error", Title: "Error", Body: strings.ReplaceAll(msg, "_", " ")}, true
}
- if flash, ok := flashFromQueryMessage(msg); ok {
+ if flash, ok := FlashFromQueryMessage(msg); ok {
return flash, true
}
return render.FlashMessage{Kind: "error", Title: "Action failed", Body: msg}, msg != ""
diff --git a/core/server/web/cookie_helpers.go b/core/server/web/cookie_helpers.go
index 5fb90959..09964536 100644
--- a/core/server/web/cookie_helpers.go
+++ b/core/server/web/cookie_helpers.go
@@ -53,3 +53,10 @@ func effectiveSessionCookieName() string {
}
const hostPrefixedSessionCookieName = "__Host-sumeru_session"
+
+func sessionCookieSecure() bool {
+ if config.AppConfig.ForceSecureCookies {
+ return true
+ }
+ return !config.AppConfig.DevMode
+}
diff --git a/core/server/web/record_error_flash_cookie.go b/core/server/web/flash_cookies.go
similarity index 50%
rename from core/server/web/record_error_flash_cookie.go
rename to core/server/web/flash_cookies.go
index 396d32a0..c95fadd7 100644
--- a/core/server/web/record_error_flash_cookie.go
+++ b/core/server/web/flash_cookies.go
@@ -6,7 +6,10 @@ import (
"net/http"
)
-const recordErrorFlashCookie = "sumeru_record_error_flash"
+const (
+ apiKeyFlashCookie = "sumeru_api_key_flash"
+ recordErrorFlashCookie = "sumeru_record_error_flash"
+)
type recordErrorFlashPayload struct {
Kind string `json:"kind"`
@@ -16,6 +19,30 @@ type recordErrorFlashPayload struct {
FieldErrors []string `json:"field_errors,omitempty"`
}
+// PageFlash is a one-time user-visible banner after redirect.
+type PageFlash struct {
+ Kind string `json:"kind"` // success, info, warning, error
+ Title string `json:"title"`
+ Body string `json:"body"`
+ Details string `json:"details,omitempty"`
+ FieldErrors []string `json:"field_errors,omitempty"`
+}
+
+// SetAPIKeyFlash stores a one-time raw API key in an HttpOnly cookie (never in redirect URLs).
+func SetAPIKeyFlash(w http.ResponseWriter, raw string) {
+ setNamedCookie(w, apiKeyFlashCookie, raw, "/", 120, http.SameSiteLaxMode)
+}
+
+// ConsumeAPIKeyFlash reads and clears the one-time API key flash cookie.
+func ConsumeAPIKeyFlash(r *http.Request, w http.ResponseWriter) string {
+ c, err := r.Cookie(apiKeyFlashCookie)
+ clearNamedCookie(w, apiKeyFlashCookie, "/", http.SameSiteLaxMode)
+ if err != nil || c.Value == "" {
+ return ""
+ }
+ return c.Value
+}
+
func recordErrorFlashCookieAttrs() *http.Cookie {
return buildNamedCookie(recordErrorFlashCookie, "", "/", 120, http.SameSiteLaxMode, true, sessionCookieSecure())
}
@@ -52,3 +79,19 @@ func ConsumeRecordErrorFlash(r *http.Request, w http.ResponseWriter) (PageFlash,
}
return PageFlash(payload), true
}
+
+// ConsumePageFlashes reads and clears one-time flash data (cookies).
+func ConsumePageFlashes(r *http.Request, w http.ResponseWriter) []PageFlash {
+ var out []PageFlash
+ if c, err := r.Cookie(apiKeyFlashCookie); err == nil && c.Value != "" {
+ out = append(out, PageFlash{
+ Kind: "success",
+ Title: "API key created",
+ Body: "Open the one-time reveal page to copy your key (available for 2 minutes):\n/web/apikey/reveal",
+ })
+ }
+ if flash, ok := ConsumeRecordErrorFlash(r, w); ok {
+ out = append(out, flash)
+ }
+ return out
+}
diff --git a/core/server/web/home_dashboard.go b/core/server/web/home_dashboard.go
index e45e6642..e96904be 100644
--- a/core/server/web/home_dashboard.go
+++ b/core/server/web/home_dashboard.go
@@ -47,11 +47,10 @@ func HomeDashboardHandler(w http.ResponseWriter, r *http.Request) {
},
MenuIDStr: "",
Page: render.PageData{
- Title: homePageTitle,
- ViewBreadcrumb: "Dashboard",
- ViewStylesheetURLs: []string{homeStylesheetURL},
- SuppressActivityDock: true,
- SuppressSidebar: true,
+ Title: homePageTitle,
+ ViewBreadcrumb: "Dashboard",
+ ViewStylesheetURLs: []string{homeStylesheetURL},
+ SuppressSidebar: true,
ViewTabs: render.HomeViewTabs(layout),
BreadcrumbItems: render.BuildHomeDashboardBreadcrumbs(ctx),
},
diff --git a/core/server/web/login.go b/core/server/web/login.go
index 218a9bad..d7d43d7d 100644
--- a/core/server/web/login.go
+++ b/core/server/web/login.go
@@ -2,11 +2,19 @@ package web
import (
"context"
+ "crypto/hmac"
+ "crypto/rand"
+ "encoding/hex"
"html/template"
"net/http"
+ "net/url"
"path/filepath"
+ "strconv"
"strings"
"sync"
+ "time"
+
+ "golang.org/x/crypto/bcrypt"
"sumeru/core/applog"
"sumeru/core/engine/assets"
@@ -14,16 +22,25 @@ import (
"sumeru/core/errcode"
"sumeru/core/mail"
"sumeru/core/orm"
+ "sumeru/core/server/auth"
"sumeru/core/server/config"
-
)
+const loginLockoutWindow = 15 * time.Minute
+
type loginPageData struct {
- Next string
- Error string
- CSRFToken string
- Stylesheets []string
- LogoURL string
+ Next string
+ Error string
+ CSRFToken string
+ Stylesheets []string
+ LogoURL string
+ AuthProviders []loginAuthProvider
+ LocalLoginEnabled bool
+ CompanyName string
+ AppName string
+ InfoTitle string
+ InfoBody string
+ Year int
}
type loginCredentials struct {
@@ -32,12 +49,361 @@ type loginCredentials struct {
Next string
}
+type loginFinishOpts struct {
+ UserID int
+ Next string
+ LoginKey string
+ ClientIP string
+ LogRoute string
+}
+
+// ponytail: in-memory lockout map; single-process only; upgrade to DB/redis for multi-instance.
var (
- loginTemplateOnce sync.Once
- cachedLoginTmpl *template.Template
- loginTemplateErr error
+ loginLockoutMu sync.Mutex
+ loginFailures = map[string][]time.Time{}
+ dummyPasswordHash = mustDummyBcryptHash()
+ authTemplateMu sync.Mutex
+ authTemplateCache = map[string]*template.Template{}
+ authTemplateErrs = map[string]error{}
)
+func mustDummyBcryptHash() string {
+ h, err := bcrypt.GenerateFromPassword([]byte("sumeru-timing-dummy-password"), bcrypt.DefaultCost)
+ if err != nil {
+ panic("login lockout: dummy bcrypt: " + err.Error())
+ }
+ return string(h)
+}
+
+func setLoginCSRFCookie(w http.ResponseWriter) string {
+ token := newLoginCSRFToken()
+ setNamedCookie(w, loginCSRFCookie, token, loginRoute, 600, http.SameSiteStrictMode)
+ return token
+}
+
+func newLoginCSRFToken() string {
+ buf := make([]byte, 16)
+ if _, err := rand.Read(buf); err != nil {
+ panic("login csrf: crypto/rand failed: " + err.Error())
+ }
+ return hex.EncodeToString(buf)
+}
+
+func validateLoginCSRF(r *http.Request) bool {
+ cookie, err := r.Cookie(loginCSRFCookie)
+ if err != nil || cookie.Value == "" {
+ return false
+ }
+ got := strings.TrimSpace(r.PostFormValue(csrfFormField))
+ if got == "" {
+ return false
+ }
+ return hmac.Equal([]byte(got), []byte(cookie.Value))
+}
+
+func clearLoginCSRFCookie(w http.ResponseWriter) {
+ clearNamedCookie(w, loginCSRFCookie, loginRoute, http.SameSiteStrictMode)
+}
+
+func setLoginNextCookie(w http.ResponseWriter, returnTo string) {
+ returnTo = SafePathNext(returnTo, homeRoute)
+ setNamedCookie(w, loginNextCookie, returnTo, loginRoute, 600, http.SameSiteLaxMode)
+}
+
+func loginNextFromRequest(r *http.Request) string {
+ cookie, err := r.Cookie(loginNextCookie)
+ if err != nil || cookie.Value == "" {
+ return ""
+ }
+ return SafePathNext(cookie.Value, homeRoute)
+}
+
+func resolveLoginNext(r *http.Request) string {
+ var next string
+ if n := loginNextFromRequest(r); n != "" {
+ next = n
+ } else if q := strings.TrimSpace(r.URL.Query().Get(nextField)); q != "" {
+ next = SafePathNext(q, homeRoute)
+ } else {
+ next = homeRoute
+ }
+ if uid := SessionUserID(r); uid > 0 {
+ return postLoginDestination(r.Context(), uid, next)
+ }
+ return next
+}
+
+func clearLoginNextCookie(w http.ResponseWriter) {
+ clearNamedCookie(w, loginNextCookie, loginRoute, http.SameSiteLaxMode)
+}
+
+func redirectToLogin(w http.ResponseWriter, r *http.Request, returnTo string) {
+ setLoginNextCookie(w, returnTo)
+ http.Redirect(w, r, loginRoute, http.StatusFound)
+}
+
+func normalizeLoginKey(login string) string {
+ return strings.ToLower(strings.TrimSpace(login))
+}
+
+func loginLocked(login string) bool {
+ key := normalizeLoginKey(login)
+ if key == "" {
+ return false
+ }
+ loginLockoutMu.Lock()
+ defer loginLockoutMu.Unlock()
+ cutoff := time.Now().Add(-loginLockoutWindow)
+ attempts := pruneAttemptsAfter(loginFailures[key], cutoff)
+ loginFailures[key] = attempts
+ return len(attempts) >= loginLockoutMaxFailures
+}
+
+func recordLoginFailure(login string) {
+ key := normalizeLoginKey(login)
+ if key == "" {
+ return
+ }
+ loginLockoutMu.Lock()
+ defer loginLockoutMu.Unlock()
+ cutoff := time.Now().Add(-loginLockoutWindow)
+ loginFailures[key] = append(pruneAttemptsAfter(loginFailures[key], cutoff), time.Now())
+}
+
+func clearLoginFailures(login string) {
+ key := normalizeLoginKey(login)
+ if key == "" {
+ return
+ }
+ loginLockoutMu.Lock()
+ delete(loginFailures, key)
+ loginLockoutMu.Unlock()
+}
+
+func resetLoginLockoutState() {
+ loginLockoutMu.Lock()
+ loginFailures = map[string][]time.Time{}
+ loginLockoutMu.Unlock()
+}
+
+func comparePasswordConstantTime(storedHash, plain string) bool {
+ storedHash = strings.TrimSpace(storedHash)
+ if storedHash == "" {
+ _ = bcrypt.CompareHashAndPassword([]byte(dummyPasswordHash), []byte(plain))
+ return false
+ }
+ return bcrypt.CompareHashAndPassword([]byte(storedHash), []byte(plain)) == nil
+}
+
+func parsePositiveInt(s string) int {
+ n, err := strconv.Atoi(strings.TrimSpace(s))
+ if err != nil || n < 0 {
+ return 0
+ }
+ return n
+}
+
+func getAuthTemplate(filename string) (*template.Template, error) {
+ authTemplateMu.Lock()
+ defer authTemplateMu.Unlock()
+ if _, loaded := authTemplateErrs[filename]; loaded {
+ return authTemplateCache[filename], authTemplateErrs[filename]
+ }
+ path := filepath.Join(config.AppConfig.TemplatesPath, filename)
+ tmpl, err := template.ParseFiles(path)
+ authTemplateCache[filename] = tmpl
+ authTemplateErrs[filename] = err
+ return tmpl, err
+}
+
+func buildLoginPageData(r *http.Request, next, errorMessage, csrfToken string) loginPageData {
+ ctx := r.Context()
+ data := loginPageData{
+ Next: next,
+ Error: errorMessage,
+ CSRFToken: csrfToken,
+ Stylesheets: assets.LoginStylesheetURLs(),
+ LogoURL: render.ShellLogoURL(),
+ AuthProviders: listEnabledAuthProviders(ctx),
+ LocalLoginEnabled: authLocalEnabled(ctx),
+ CompanyName: loginPageCompanyName(ctx),
+ AppName: "Sumeru",
+ Year: time.Now().Year(),
+ }
+ if flash, ok := FlashFromQueryMessage(r.URL.Query().Get(flashMessageParam)); ok {
+ if errorMessage == "" && flash.Kind == "error" {
+ data.Error = flash.Body
+ if flash.Title != "" {
+ data.Error = flash.Title + ": " + flash.Body
+ }
+ } else if flash.Body != "" {
+ data.InfoTitle = flash.Title
+ data.InfoBody = flash.Body
+ }
+ }
+ return data
+}
+
+func writeAuthFormPage(w http.ResponseWriter, r *http.Request, templateFile, logRoute string, statusCode int, next, errorMessage, csrfToken string) {
+ tmpl, err := getAuthTemplate(templateFile)
+ if err != nil {
+ if statusCode == http.StatusOK {
+ WebLogEvent(r.Context(), WebLogInput{
+ Route: logRoute,
+ Message: "auth form template unavailable",
+ Code: errcode.InternalError,
+ Operation: "login_template",
+ Status: logStatusFailure,
+ Err: err,
+ })
+ http.Error(w, "Page unavailable", http.StatusInternalServerError)
+ return
+ }
+ http.Error(w, errorMessage, http.StatusUnauthorized)
+ return
+ }
+
+ w.Header().Set("Content-Type", "text/html; charset=utf-8")
+ if statusCode != http.StatusOK {
+ w.WriteHeader(statusCode)
+ }
+ _ = tmpl.Execute(w, buildLoginPageData(r, next, errorMessage, csrfToken))
+}
+
+func showAuthFormError(w http.ResponseWriter, r *http.Request, templateFile, logRoute string, statusCode int, next, errorMessage string) {
+ csrfToken := setLoginCSRFCookie(w)
+ writeAuthFormPage(w, r, templateFile, logRoute, statusCode, next, errorMessage, csrfToken)
+}
+
+func registerTOTPRoutes() {
+ registerPublic(http.MethodGet, totpLoginRoute, TOTPLoginGet)
+ registerPublic(http.MethodPost, totpLoginRoute, TOTPLoginPost)
+}
+
+func TOTPLoginGet(w http.ResponseWriter, r *http.Request) {
+ if pendingMFAUserID(r) <= 0 {
+ http.Redirect(w, r, loginRoute, http.StatusFound)
+ return
+ }
+ csrfToken := setLoginCSRFCookie(w)
+ writeAuthFormPage(w, r, totpLoginTemplateFile, totpLoginRoute, http.StatusOK, resolveLoginNext(r), "", csrfToken)
+}
+
+func TOTPLoginPost(w http.ResponseWriter, r *http.Request) {
+ if !ParsePostForm(w, r) {
+ return
+ }
+ uid := pendingMFAUserID(r)
+ if uid <= 0 {
+ http.Error(w, "session expired", http.StatusUnauthorized)
+ return
+ }
+ next := SafePathNext(r.PostFormValue(nextField), homeRoute)
+ if !validateLoginCSRF(r) {
+ showAuthFormError(w, r, totpLoginTemplateFile, totpLoginRoute, http.StatusForbidden, next, "Invalid form")
+ return
+ }
+ code := strings.TrimSpace(r.PostFormValue("totp_code"))
+ secret, _ := userTOTPFields(r.Context(), uid)
+ if !auth.ValidateTOTP(secret, code) {
+ showAuthFormError(w, r, totpLoginTemplateFile, totpLoginRoute, http.StatusUnauthorized, next, "Invalid authentication code")
+ return
+ }
+ if strings.TrimSpace(r.PostFormValue("trust_device")) == "1" {
+ setTrustedDeviceCookie(w, uid)
+ }
+ finishTOTPSession(w, r, uid, next)
+}
+
+func userTOTPFields(ctx context.Context, uid int) (secret string, enabled bool) {
+ return orm.UserTOTPCredentialsForLogin(ctx, uid)
+}
+
+func needsTOTP(r *http.Request, uid int) bool {
+ if trustedDeviceValid(r, uid) {
+ return false
+ }
+ _, enabled := userTOTPFields(r.Context(), uid)
+ return enabled
+}
+
+func trustedDeviceValid(r *http.Request, expectUID int) bool {
+ c, err := r.Cookie(trustedDeviceCookie)
+ if err != nil || strings.TrimSpace(c.Value) == "" {
+ return false
+ }
+ cookieUID, ok := parseSignedUIDCookie(c.Value)
+ return ok && cookieUID == expectUID
+}
+
+func setPendingMFACookie(w http.ResponseWriter, uid int) {
+ setNamedCookie(w, pendingMFACookie, mintSignedUIDCookieValue(uid, 5*time.Minute), "/", 300, sessionSameSite())
+}
+
+func pendingMFAUserID(r *http.Request) int {
+ c, err := r.Cookie(pendingMFACookie)
+ if err != nil {
+ return 0
+ }
+ uid, ok := parseSignedUIDCookie(c.Value)
+ if !ok {
+ return 0
+ }
+ return uid
+}
+
+func clearPendingMFACookie(w http.ResponseWriter) {
+ clearNamedCookie(w, pendingMFACookie, "/", sessionSameSite())
+}
+
+func setTrustedDeviceCookie(w http.ResponseWriter, uid int) {
+ setNamedCookie(w, trustedDeviceCookie, mintSignedUIDCookieValue(uid, 30*24*time.Hour), "/", 30*24*3600, sessionSameSite())
+}
+
+func completePasswordOrOAuthLogin(w http.ResponseWriter, r *http.Request, opts loginFinishOpts) {
+ next := SafePathNext(opts.Next, homeRoute)
+ if needsTOTP(r, opts.UserID) {
+ setPendingMFACookie(w, opts.UserID)
+ if opts.LoginKey != "" {
+ clearLoginFailures(opts.LoginKey)
+ }
+ clearLoginCSRFCookie(w)
+ http.Redirect(w, r, totpLoginRoute+"?next="+url.QueryEscape(next), http.StatusSeeOther)
+ return
+ }
+ if opts.LoginKey != "" {
+ clearLoginFailures(opts.LoginKey)
+ }
+ if opts.ClientIP != "" {
+ orm.AppendUserLog(r.Context(), opts.UserID, opts.ClientIP, "success")
+ }
+ establishSessionAndRedirect(w, r, opts.UserID, next, opts.LogRoute)
+}
+
+func finishTOTPSession(w http.ResponseWriter, r *http.Request, uid int, nextFromForm string) {
+ clearPendingMFACookie(w)
+ establishSessionAndRedirect(w, r, uid, SafePathNext(nextFromForm, homeRoute), totpLoginRoute)
+}
+
+func establishSessionAndRedirect(w http.ResponseWriter, r *http.Request, userID int, next, logRoute string) {
+ if err := CreateSession(w, userID); err != nil {
+ WebLogEvent(r.Context(), WebLogInput{
+ Route: logRoute,
+ Message: "Could not start session",
+ Code: errcode.InternalError,
+ Operation: "session_create",
+ Status: logStatusFailure,
+ Err: err,
+ })
+ http.Error(w, "Could not start session", http.StatusInternalServerError)
+ return
+ }
+ clearLoginCSRFCookie(w)
+ clearLoginNextCookie(w)
+ dest := postLoginDestination(r.Context(), userID, next)
+ http.Redirect(w, r, dest, http.StatusSeeOther)
+}
+
func LoginGet(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
@@ -55,7 +421,7 @@ func LoginGet(w http.ResponseWriter, r *http.Request) {
}
csrfToken := setLoginCSRFCookie(w)
- writeLoginPage(w, r, http.StatusOK, resolveLoginNext(r), "", csrfToken)
+ writeAuthFormPage(w, r, loginTemplateFile, loginRoute, http.StatusOK, resolveLoginNext(r), "", csrfToken)
}
func LoginPost(w http.ResponseWriter, r *http.Request) {
@@ -66,16 +432,20 @@ func LoginPost(w http.ResponseWriter, r *http.Request) {
if !ParsePostForm(w, r) {
return
}
+ if !authLocalEnabled(r.Context()) {
+ showAuthFormError(w, r, loginTemplateFile, loginRoute, http.StatusForbidden,
+ SafePathNext(r.PostFormValue(nextField), homeRoute), "Password sign-in is disabled. Use single sign-on.")
+ return
+ }
if !validateLoginCSRF(r) {
- csrfToken := setLoginCSRFCookie(w)
- writeLoginPage(w, r, http.StatusForbidden, SafePathNext(r.PostFormValue(nextField), homeRoute), "Invalid or expired login form", csrfToken)
+ showAuthFormError(w, r, loginTemplateFile, loginRoute, http.StatusForbidden,
+ SafePathNext(r.PostFormValue(nextField), homeRoute), "Invalid or expired login form")
return
}
credentials := parseLoginCredentials(r)
if loginLocked(credentials.Login) {
- csrfToken := setLoginCSRFCookie(w)
- writeLoginPage(w, r, http.StatusUnauthorized, credentials.Next, invalidLoginMessage, csrfToken)
+ showAuthFormError(w, r, loginTemplateFile, loginRoute, http.StatusUnauthorized, credentials.Next, invalidLoginMessage)
return
}
clientIP := clientIP(r)
@@ -91,29 +461,16 @@ func LoginPost(w http.ResponseWriter, r *http.Request) {
"ip": clientIP,
},
})
- csrfToken := setLoginCSRFCookie(w)
- writeLoginPage(w, r, http.StatusUnauthorized, credentials.Next, invalidLoginMessage, csrfToken)
+ showAuthFormError(w, r, loginTemplateFile, loginRoute, http.StatusUnauthorized, credentials.Next, invalidLoginMessage)
return
}
- if err := CreateSession(w, userID); err != nil {
- WebLogEvent(r.Context(), WebLogInput{
- Route: loginRoute,
- Message: "Could not start session",
- Code: errcode.InternalError,
- Operation: "session_create",
- Status: logStatusFailure,
- Err: err,
- })
- http.Error(w, "Could not start session", http.StatusInternalServerError)
- return
- }
-
- clearLoginFailures(credentials.Login)
- orm.AppendUserLog(r.Context(), userID, clientIP, "success")
- clearLoginCSRFCookie(w)
- clearLoginNextCookie(w)
- dest := postLoginDestination(r.Context(), userID, credentials.Next)
- http.Redirect(w, r, dest, http.StatusSeeOther)
+ completePasswordOrOAuthLogin(w, r, loginFinishOpts{
+ UserID: userID,
+ Next: credentials.Next,
+ LoginKey: credentials.Login,
+ ClientIP: clientIP,
+ LogRoute: loginRoute,
+ })
}
func LogoutGet(w http.ResponseWriter, r *http.Request) {
@@ -178,46 +535,6 @@ func recordFailedLogin(ctx context.Context, userID int, clientIP, auditNote stri
orm.AppendAudit(ctx, "login_fail", coreUserModel, int64(userID), nil, nil, auditNote)
}
-func getLoginTemplate() (*template.Template, error) {
- loginTemplateOnce.Do(func() {
- templatePath := filepath.Join(config.AppConfig.TemplatesPath, loginTemplateFile)
- cachedLoginTmpl, loginTemplateErr = template.ParseFiles(templatePath)
- })
- return cachedLoginTmpl, loginTemplateErr
-}
-
-func writeLoginPage(w http.ResponseWriter, r *http.Request, statusCode int, next, errorMessage, csrfToken string) {
- tmpl, err := getLoginTemplate()
- if err != nil {
- if statusCode == http.StatusOK {
- WebLogEvent(r.Context(), WebLogInput{
- Route: loginRoute,
- Message: "login template unavailable",
- Code: errcode.InternalError,
- Operation: "login_template",
- Status: logStatusFailure,
- Err: err,
- })
- http.Error(w, "Login page unavailable", http.StatusInternalServerError)
- return
- }
- http.Error(w, errorMessage, http.StatusUnauthorized)
- return
- }
-
- w.Header().Set("Content-Type", "text/html; charset=utf-8")
- if statusCode != http.StatusOK {
- w.WriteHeader(statusCode)
- }
- _ = tmpl.Execute(w, loginPageData{
- Next: next,
- Error: errorMessage,
- CSRFToken: csrfToken,
- Stylesheets: assets.LoginStylesheetURLs(),
- LogoURL: render.ShellLogoURL(),
- })
-}
-
func ActionResetPassword(w http.ResponseWriter, r *http.Request) {
if !requireLoginAndPOST(w, r) {
return
diff --git a/core/server/web/login_csrf.go b/core/server/web/login_csrf.go
deleted file mode 100644
index 617099a1..00000000
--- a/core/server/web/login_csrf.go
+++ /dev/null
@@ -1,50 +0,0 @@
-package web
-
-import (
- "crypto/hmac"
- "crypto/rand"
- "encoding/hex"
- "net/http"
- "strings"
-
- "sumeru/core/server/config"
-)
-
-const loginCSRFCookie = "sumeru_login_csrf"
-
-func setLoginCSRFCookie(w http.ResponseWriter) string {
- token := newLoginCSRFToken()
- setNamedCookie(w, loginCSRFCookie, token, loginRoute, 600, http.SameSiteStrictMode)
- return token
-}
-
-func newLoginCSRFToken() string {
- buf := make([]byte, 16)
- if _, err := rand.Read(buf); err != nil {
- panic("login csrf: crypto/rand failed: " + err.Error())
- }
- return hex.EncodeToString(buf)
-}
-
-func validateLoginCSRF(r *http.Request) bool {
- cookie, err := r.Cookie(loginCSRFCookie)
- if err != nil || cookie.Value == "" {
- return false
- }
- got := strings.TrimSpace(r.PostFormValue(csrfFormField))
- if got == "" {
- return false
- }
- return hmac.Equal([]byte(got), []byte(cookie.Value))
-}
-
-func clearLoginCSRFCookie(w http.ResponseWriter) {
- clearNamedCookie(w, loginCSRFCookie, loginRoute, http.SameSiteStrictMode)
-}
-
-func sessionCookieSecure() bool {
- if config.AppConfig.ForceSecureCookies {
- return true
- }
- return !config.AppConfig.DevMode
-}
diff --git a/core/server/web/login_lockout.go b/core/server/web/login_lockout.go
deleted file mode 100644
index 98ca07e7..00000000
--- a/core/server/web/login_lockout.go
+++ /dev/null
@@ -1,84 +0,0 @@
-package web
-
-import (
- "strings"
- "sync"
- "time"
-
- "golang.org/x/crypto/bcrypt"
-)
-
-// ponytail: in-memory lockout map; single-process only; upgrade to DB/redis for multi-instance.
-var (
- loginLockoutMu sync.Mutex
- loginFailures = map[string][]time.Time{}
-)
-
-const (
- loginLockoutMaxFailures = 5
- loginLockoutWindow = 15 * time.Minute
-)
-
-// bcrypt dummy hash (cost 10) for constant-time path on unknown login.
-var dummyPasswordHash = mustDummyBcryptHash()
-
-func mustDummyBcryptHash() string {
- h, err := bcrypt.GenerateFromPassword([]byte("sumeru-timing-dummy-password"), bcrypt.DefaultCost)
- if err != nil {
- panic("login lockout: dummy bcrypt: " + err.Error())
- }
- return string(h)
-}
-
-func normalizeLoginKey(login string) string {
- return strings.ToLower(strings.TrimSpace(login))
-}
-
-func loginLocked(login string) bool {
- key := normalizeLoginKey(login)
- if key == "" {
- return false
- }
- loginLockoutMu.Lock()
- defer loginLockoutMu.Unlock()
- cutoff := time.Now().Add(-loginLockoutWindow)
- attempts := pruneAttemptsAfter(loginFailures[key], cutoff)
- loginFailures[key] = attempts
- return len(attempts) >= loginLockoutMaxFailures
-}
-
-func recordLoginFailure(login string) {
- key := normalizeLoginKey(login)
- if key == "" {
- return
- }
- loginLockoutMu.Lock()
- defer loginLockoutMu.Unlock()
- cutoff := time.Now().Add(-loginLockoutWindow)
- loginFailures[key] = append(pruneAttemptsAfter(loginFailures[key], cutoff), time.Now())
-}
-
-func clearLoginFailures(login string) {
- key := normalizeLoginKey(login)
- if key == "" {
- return
- }
- loginLockoutMu.Lock()
- delete(loginFailures, key)
- loginLockoutMu.Unlock()
-}
-
-func resetLoginLockoutState() {
- loginLockoutMu.Lock()
- loginFailures = map[string][]time.Time{}
- loginLockoutMu.Unlock()
-}
-
-func comparePasswordConstantTime(storedHash, plain string) bool {
- storedHash = strings.TrimSpace(storedHash)
- if storedHash == "" {
- _ = bcrypt.CompareHashAndPassword([]byte(dummyPasswordHash), []byte(plain))
- return false
- }
- return bcrypt.CompareHashAndPassword([]byte(storedHash), []byte(plain)) == nil
-}
diff --git a/core/server/web/login_next.go b/core/server/web/login_next.go
deleted file mode 100644
index 309e4ed7..00000000
--- a/core/server/web/login_next.go
+++ /dev/null
@@ -1,45 +0,0 @@
-package web
-
-import (
- "net/http"
- "strings"
-)
-
-const loginNextCookie = "sumeru_login_next"
-
-func setLoginNextCookie(w http.ResponseWriter, returnTo string) {
- returnTo = SafePathNext(returnTo, homeRoute)
- setNamedCookie(w, loginNextCookie, returnTo, loginRoute, 600, http.SameSiteLaxMode)
-}
-
-func loginNextFromRequest(r *http.Request) string {
- cookie, err := r.Cookie(loginNextCookie)
- if err != nil || cookie.Value == "" {
- return ""
- }
- return SafePathNext(cookie.Value, homeRoute)
-}
-
-func resolveLoginNext(r *http.Request) string {
- var next string
- if n := loginNextFromRequest(r); n != "" {
- next = n
- } else if q := strings.TrimSpace(r.URL.Query().Get(nextField)); q != "" {
- next = SafePathNext(q, homeRoute)
- } else {
- next = homeRoute
- }
- if uid := SessionUserID(r); uid > 0 {
- return postLoginDestination(r.Context(), uid, next)
- }
- return next
-}
-
-func clearLoginNextCookie(w http.ResponseWriter) {
- clearNamedCookie(w, loginNextCookie, loginRoute, http.SameSiteLaxMode)
-}
-
-func redirectToLogin(w http.ResponseWriter, r *http.Request, returnTo string) {
- setLoginNextCookie(w, returnTo)
- http.Redirect(w, r, loginRoute, http.StatusFound)
-}
diff --git a/core/server/web/oauth_login.go b/core/server/web/oauth_login.go
new file mode 100644
index 00000000..908bda4e
--- /dev/null
+++ b/core/server/web/oauth_login.go
@@ -0,0 +1,406 @@
+package web
+
+import (
+ "context"
+ "crypto/rand"
+ "crypto/sha256"
+ "encoding/base64"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "net/url"
+ "strings"
+ "time"
+
+ "sumeru/core/orm"
+ "sumeru/core/server/auth"
+ "sumeru/core/server/config"
+)
+
+const (
+ oauthStartRoute = "/web/auth/oauth/start"
+ oauthCallbackRoute = "/web/auth/oauth/callback"
+ oauthStateCookie = "sumeru_oauth_state"
+)
+
+func registerOAuthRoutes() {
+ registerPublic(http.MethodGet, oauthStartRoute, OAuthStartHandler)
+ registerPublic(http.MethodGet, oauthCallbackRoute, OAuthCallbackHandler)
+}
+
+func OAuthStartHandler(w http.ResponseWriter, r *http.Request) {
+ providerID := parsePositiveInt(r.URL.Query().Get("provider"))
+ if providerID <= 0 {
+ http.Error(w, "provider required", http.StatusBadRequest)
+ return
+ }
+ ctx := orm.AuditedBypass(r.Context(), "oauth.start")
+ prov, err := orm.SearchOne(ctx, "sys.auth.provider", map[string]interface{}{"id": providerID, "enabled": true})
+ if err != nil || prov == nil {
+ http.Error(w, "provider not found", http.StatusNotFound)
+ return
+ }
+ state := randomURLSafe(32)
+ verifier := randomURLSafe(64)
+ challenge := pkceChallenge(verifier)
+ setOAuthStateCookie(w, state, verifier, providerID)
+ authURL := strings.TrimSpace(oauthRowString(prov, "authorize_url"))
+ if authURL == "" {
+ authURL = strings.TrimRight(strings.TrimSpace(oauthRowString(prov, "issuer_url")), "/") + "/authorize"
+ }
+ q := url.Values{}
+ q.Set("response_type", "code")
+ q.Set("client_id", oauthRowString(prov, "client_id"))
+ q.Set("redirect_uri", oauthRedirectURI(r))
+ q.Set("scope", defaultScope(oauthRowString(prov, "scopes")))
+ q.Set("state", state)
+ q.Set("code_challenge", challenge)
+ q.Set("code_challenge_method", "S256")
+ sep := "?"
+ if strings.Contains(authURL, "?") {
+ sep = "&"
+ }
+ http.Redirect(w, r, authURL+sep+q.Encode(), http.StatusFound)
+}
+
+func OAuthCallbackHandler(w http.ResponseWriter, r *http.Request) {
+ state := r.URL.Query().Get("state")
+ code := r.URL.Query().Get("code")
+ providerID, verifier, ok := validateOAuthStateCookie(r, state)
+ if !ok || code == "" {
+ redirectLoginOAuthError(w, r)
+ return
+ }
+ clearOAuthStateCookie(w)
+ ctx := r.Context()
+ prov, err := orm.SearchOne(orm.AuditedBypass(ctx, "oauth.callback"), "sys.auth.provider", map[string]interface{}{"id": providerID})
+ if err != nil || prov == nil {
+ redirectLoginOAuthError(w, r)
+ return
+ }
+ tokenURL := strings.TrimSpace(oauthRowString(prov, "token_url"))
+ if tokenURL == "" {
+ tokenURL = strings.TrimRight(strings.TrimSpace(oauthRowString(prov, "issuer_url")), "/") + "/token"
+ }
+ tok, err := exchangeOAuthCode(r, tokenURL, oauthRowString(prov, "client_id"), oauthRowString(prov, "client_secret"), code, verifier, oauthRedirectURI(r))
+ if err != nil {
+ redirectLoginOAuthError(w, r)
+ return
+ }
+ email, sub, err := resolveOAuthIdentity(r.Context(), prov, tok)
+ if err != nil {
+ redirectLoginOAuthError(w, r)
+ return
+ }
+ if err := validateOAuthIdentityClaims(prov, tok, sub); err != nil {
+ redirectLoginOAuthError(w, r)
+ return
+ }
+ userID, err := linkOAuthUser(ctx, providerID, sub, email, oauthRowString(prov, "link_policy"))
+ if err != nil || userID <= 0 {
+ redirectLoginOAuthError(w, r)
+ return
+ }
+ completePasswordOrOAuthLogin(w, r, loginFinishOpts{
+ UserID: userID,
+ Next: resolveLoginNext(r),
+ LogRoute: oauthCallbackRoute,
+ })
+}
+
+func exchangeOAuthCode(r *http.Request, tokenURL, clientID, clientSecret, code, verifier, redirectURI string) (map[string]interface{}, error) {
+ form := url.Values{}
+ form.Set("grant_type", "authorization_code")
+ form.Set("code", code)
+ form.Set("redirect_uri", redirectURI)
+ form.Set("client_id", clientID)
+ form.Set("code_verifier", verifier)
+ if clientSecret != "" {
+ form.Set("client_secret", clientSecret)
+ }
+ req, err := http.NewRequestWithContext(r.Context(), http.MethodPost, tokenURL, strings.NewReader(form.Encode()))
+ if err != nil {
+ return nil, err
+ }
+ req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
+ resp, err := http.DefaultClient.Do(req)
+ if err != nil {
+ return nil, err
+ }
+ defer resp.Body.Close()
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
+ if resp.StatusCode >= 400 {
+ return nil, fmt.Errorf("token http %d", resp.StatusCode)
+ }
+ var out map[string]interface{}
+ if err := json.Unmarshal(body, &out); err != nil {
+ return nil, err
+ }
+ return out, nil
+}
+
+func linkOAuthUser(ctx context.Context, providerID int, subject, email, policy string) (int, error) {
+ subject = strings.TrimSpace(subject)
+ if subject == "" {
+ return 0, fmt.Errorf("missing subject")
+ }
+ if rec, err := orm.SearchOne(ctx, "core.user.identity", map[string]interface{}{
+ "provider_id": providerID,
+ "subject": subject,
+ }); err == nil && rec != nil {
+ return int(oauthRowInt(rec, "user_id")), nil
+ }
+ email = strings.TrimSpace(strings.ToLower(email))
+ if email != "" {
+ if user, err := orm.SearchOne(ctx, "core.user", map[string]interface{}{"login": email}); err == nil && user != nil {
+ uid := int(oauthRowInt(user, "id"))
+ _ = upsertIdentity(ctx, providerID, subject, uid)
+ return uid, nil
+ }
+ }
+ switch strings.TrimSpace(policy) {
+ case "jit_internal", "jit_portal":
+ // ponytail: minimal JIT — create inactive-safe internal user with login=email or subject.
+ login := email
+ if login == "" {
+ login = "oidc_" + subject
+ }
+ inst := orm.Registry["core.user"]
+ uid, err := orm.Create(ctx, inst, map[string]interface{}{
+ "login": login,
+ "name": login,
+ "email": email,
+ "active": true,
+ "user_type": map[bool]string{true: "portal", false: "internal"}[policy == "jit_portal"],
+ })
+ if err != nil {
+ return 0, err
+ }
+ _ = upsertIdentity(ctx, providerID, subject, uid)
+ return uid, nil
+ default:
+ return 0, fmt.Errorf("unknown user")
+ }
+}
+
+func upsertIdentity(ctx context.Context, providerID int, subject string, uid int) error {
+ inst := orm.Registry["core.user.identity"]
+ _, err := orm.Create(ctx, inst, map[string]interface{}{
+ "provider_id": providerID,
+ "subject": subject,
+ "user_id": uid,
+ })
+ return err
+}
+
+func validateOAuthIdentityClaims(prov map[string]interface{}, tok map[string]interface{}, sub string) error {
+ sub = strings.TrimSpace(sub)
+ if sub == "" {
+ return fmt.Errorf("missing subject")
+ }
+ if !config.AppConfig.DevMode {
+ if strings.TrimSpace(oauthRowString(prov, "jwks_url")) != "" {
+ idTok, _ := tok["id_token"].(string)
+ if strings.TrimSpace(idTok) == "" {
+ return fmt.Errorf("id_token required")
+ }
+ }
+ }
+ return nil
+}
+
+func resolveOAuthIdentity(ctx context.Context, prov map[string]interface{}, tok map[string]interface{}) (email, sub string, err error) {
+ idTok, _ := tok["id_token"].(string)
+ jwksURL := strings.TrimSpace(oauthRowString(prov, "jwks_url"))
+ issuer := strings.TrimSpace(oauthRowString(prov, "issuer_url"))
+ clientID := strings.TrimSpace(oauthRowString(prov, "client_id"))
+ if strings.TrimSpace(idTok) != "" && jwksURL != "" {
+ claims, verr := auth.VerifyIDToken(ctx, idTok, issuer, clientID, jwksURL)
+ if verr != nil {
+ return "", "", verr
+ }
+ email, _ = claims["email"].(string)
+ sub, _ = claims["sub"].(string)
+ return email, sub, nil
+ }
+ if strings.TrimSpace(idTok) != "" {
+ email, sub, _ = auth.ParseIDTokenClaimsUnsafe(idTok)
+ }
+ if sub == "" {
+ sub, _ = tok["sub"].(string)
+ }
+ if email == "" {
+ email, _ = tok["email"].(string)
+ }
+ if sub == "" && strings.TrimSpace(oauthRowString(prov, "provider_type")) == "oauth2" {
+ if access, _ := tok["access_token"].(string); strings.TrimSpace(access) != "" {
+ email, sub = fetchOAuthUserInfo(ctx, access, issuer)
+ }
+ }
+ if sub == "" {
+ return "", "", fmt.Errorf("missing subject")
+ }
+ return email, sub, nil
+}
+
+func fetchOAuthUserInfo(ctx context.Context, accessToken, issuer string) (email, sub string) {
+ base := strings.TrimRight(strings.TrimSpace(issuer), "/")
+ if base == "" {
+ return "", ""
+ }
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, base+"/userinfo", nil)
+ if err != nil {
+ return "", ""
+ }
+ req.Header.Set("Authorization", "Bearer "+accessToken)
+ resp, err := http.DefaultClient.Do(req)
+ if err != nil {
+ return "", ""
+ }
+ defer resp.Body.Close()
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
+ if resp.StatusCode >= 400 {
+ return "", ""
+ }
+ var claims map[string]interface{}
+ if json.Unmarshal(body, &claims) != nil {
+ return "", ""
+ }
+ email, _ = claims["email"].(string)
+ sub, _ = claims["sub"].(string)
+ return email, sub
+}
+
+func redirectLoginOAuthError(w http.ResponseWriter, r *http.Request) {
+ http.Redirect(w, r, loginRoute+"?"+flashMessageParam+"="+url.QueryEscape(oauthDeniedMsg), http.StatusFound)
+}
+
+type loginAuthProvider struct {
+ ID int
+ Label string
+}
+
+func listEnabledAuthProviders(ctx context.Context) []loginAuthProvider {
+ ctx = orm.AuditedBypass(ctx, "login.providers")
+ rows, err := orm.Search(ctx, "sys.auth.provider", [][]interface{}{{"enabled", "=", true}})
+ if err != nil || len(rows) == 0 {
+ return nil
+ }
+ out := make([]loginAuthProvider, 0, len(rows))
+ for _, row := range rows {
+ id := int(oauthRowInt(row, "id"))
+ if id <= 0 {
+ continue
+ }
+ label := strings.TrimSpace(oauthRowString(row, "button_label"))
+ if label == "" {
+ label = strings.TrimSpace(oauthRowString(row, "name"))
+ }
+ if label == "" {
+ label = "Sign in"
+ }
+ out = append(out, loginAuthProvider{ID: id, Label: label})
+ }
+ return out
+}
+
+func authLocalEnabled(ctx context.Context) bool {
+ raw := strings.TrimSpace(strings.ToLower(orm.GetConfig(orm.AuditedBypass(ctx, "login.config"), authLocalConfigKey, "true")))
+ switch raw {
+ case "0", "false", "no", "off":
+ providers := listEnabledAuthProviders(ctx)
+ return len(providers) == 0
+ default:
+ return true
+ }
+}
+
+func loginPageCompanyName(ctx context.Context) string {
+ ctx = orm.AuditedBypass(ctx, "login.brand")
+ rows, err := orm.SearchLimit(ctx, "core.company", nil, 1)
+ if err != nil || len(rows) == 0 {
+ return ""
+ }
+ return strings.TrimSpace(oauthRowString(rows[0], "name"))
+}
+
+func pkceChallenge(verifier string) string {
+ sum := sha256.Sum256([]byte(verifier))
+ return base64.RawURLEncoding.EncodeToString(sum[:])
+}
+
+func randomURLSafe(n int) string {
+ b := make([]byte, n)
+ _, _ = rand.Read(b)
+ return base64.RawURLEncoding.EncodeToString(b)
+}
+
+func oauthRedirectURI(r *http.Request) string {
+ scheme := "http"
+ if r.TLS != nil || r.Header.Get("X-Forwarded-Proto") == "https" {
+ scheme = "https"
+ }
+ return scheme + "://" + r.Host + oauthCallbackRoute
+}
+
+func defaultScope(s string) string {
+ s = strings.TrimSpace(s)
+ if s == "" {
+ return "openid email profile"
+ }
+ return s
+}
+
+func oauthRowString(rec map[string]interface{}, key string) string {
+ if v, ok := rec[key].(string); ok {
+ return v
+ }
+ return ""
+}
+
+func oauthRowInt(rec map[string]interface{}, key string) int64 {
+ switch v := rec[key].(type) {
+ case int:
+ return int64(v)
+ case int64:
+ return v
+ case float64:
+ return int64(v)
+ default:
+ return 0
+ }
+}
+
+// OAuth state cookie helpers (state|verifier|providerID base64 json).
+func setOAuthStateCookie(w http.ResponseWriter, state, verifier string, providerID int) {
+ payload, _ := json.Marshal(map[string]interface{}{
+ "state": state, "verifier": verifier, "provider_id": providerID, "exp": time.Now().Add(10 * time.Minute).Unix(),
+ })
+ setNamedCookie(w, oauthStateCookie, base64.RawURLEncoding.EncodeToString(payload), "/", 600, sessionSameSite())
+}
+
+func validateOAuthStateCookie(r *http.Request, state string) (providerID int, verifier string, ok bool) {
+ c, err := r.Cookie(oauthStateCookie)
+ if err != nil || c.Value == "" {
+ return 0, "", false
+ }
+ raw, err := base64.RawURLEncoding.DecodeString(c.Value)
+ if err != nil {
+ return 0, "", false
+ }
+ var m map[string]interface{}
+ if json.Unmarshal(raw, &m) != nil {
+ return 0, "", false
+ }
+ if fmt.Sprint(m["state"]) != state {
+ return 0, "", false
+ }
+ providerID = int(oauthRowInt(m, "provider_id"))
+ verifier, _ = m["verifier"].(string)
+ return providerID, verifier, providerID > 0 && verifier != ""
+}
+
+func clearOAuthStateCookie(w http.ResponseWriter) {
+ clearNamedCookie(w, oauthStateCookie, "/", sessionSameSite())
+}
diff --git a/core/server/web/page_flash.go b/core/server/web/page_flash.go
deleted file mode 100644
index 20d23b7d..00000000
--- a/core/server/web/page_flash.go
+++ /dev/null
@@ -1,30 +0,0 @@
-package web
-
-import (
- "net/http"
-)
-
-// PageFlash is a one-time user-visible banner after redirect.
-type PageFlash struct {
- Kind string `json:"kind"` // success, info, warning, error
- Title string `json:"title"`
- Body string `json:"body"`
- Details string `json:"details,omitempty"` // optional technical details for error flashes
- FieldErrors []string `json:"field_errors,omitempty"`
-}
-
-// ConsumePageFlashes reads and clears one-time flash data (cookies).
-func ConsumePageFlashes(r *http.Request, w http.ResponseWriter) []PageFlash {
- var out []PageFlash
- if c, err := r.Cookie(apiKeyFlashCookie); err == nil && c.Value != "" {
- out = append(out, PageFlash{
- Kind: "success",
- Title: "API key created",
- Body: "Open the one-time reveal page to copy your key (available for 2 minutes):\n/web/apikey/reveal",
- })
- }
- if flash, ok := ConsumeRecordErrorFlash(r, w); ok {
- out = append(out, flash)
- }
- return out
-}
diff --git a/core/server/web/portal_auth.go b/core/server/web/portal_auth.go
deleted file mode 100644
index ea77c755..00000000
--- a/core/server/web/portal_auth.go
+++ /dev/null
@@ -1,45 +0,0 @@
-package web
-
-import (
- "context"
- "net/http"
-
- "sumeru/core/orm"
-)
-
-func requirePortalUser(w http.ResponseWriter, r *http.Request) bool {
- if !requireLogin(w, r) {
- return false
- }
- uid := SessionUserID(r)
- if uid <= 0 {
- return false
- }
- if !orm.UserIsPortalOnly(r.Context(), uid) && !orm.UserHasGroupXML(r.Context(), uid, "base.group_user") {
- http.Error(w, "Forbidden", http.StatusForbidden)
- return false
- }
- return true
-}
-
-func redirectPortalUserFromWeb(w http.ResponseWriter, r *http.Request, uid int) bool {
- if uid <= 0 || !orm.UserIsPortalOnly(r.Context(), uid) {
- return false
- }
- http.Redirect(w, r, portalHomeRoute, http.StatusFound)
- return true
-}
-
-func postLoginDestination(ctx context.Context, userID int, requestedNext string) string {
- if orm.UserIsPortalOnly(ctx, userID) {
- if stringsHasPortalPrefix(requestedNext) {
- return SafePathNext(requestedNext, portalHomeRoute)
- }
- return portalHomeRoute
- }
- return SafePathNext(requestedNext, homeRoute)
-}
-
-func stringsHasPortalPrefix(path string) bool {
- return len(path) >= len(portalRoutePrefix) && path[:len(portalRoutePrefix)] == portalRoutePrefix
-}
diff --git a/core/server/web/portal_handlers.go b/core/server/web/portal_handlers.go
index 51df60dc..3d7db816 100644
--- a/core/server/web/portal_handlers.go
+++ b/core/server/web/portal_handlers.go
@@ -1,6 +1,7 @@
package web
import (
+ "context"
"html/template"
"net/http"
"path/filepath"
@@ -182,3 +183,40 @@ func writePortalShare(w http.ResponseWriter, data portalShareData) {
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_ = shareTmpl.Execute(w, data)
}
+
+func requirePortalUser(w http.ResponseWriter, r *http.Request) bool {
+ if !requireLogin(w, r) {
+ return false
+ }
+ uid := SessionUserID(r)
+ if uid <= 0 {
+ return false
+ }
+ if !orm.UserIsPortalOnly(r.Context(), uid) && !orm.UserHasGroupXML(r.Context(), uid, "base.group_user") {
+ http.Error(w, "Forbidden", http.StatusForbidden)
+ return false
+ }
+ return true
+}
+
+func redirectPortalUserFromWeb(w http.ResponseWriter, r *http.Request, uid int) bool {
+ if uid <= 0 || !orm.UserIsPortalOnly(r.Context(), uid) {
+ return false
+ }
+ http.Redirect(w, r, portalHomeRoute, http.StatusFound)
+ return true
+}
+
+func postLoginDestination(ctx context.Context, userID int, requestedNext string) string {
+ if orm.UserIsPortalOnly(ctx, userID) {
+ if stringsHasPortalPrefix(requestedNext) {
+ return SafePathNext(requestedNext, portalHomeRoute)
+ }
+ return portalHomeRoute
+ }
+ return SafePathNext(requestedNext, homeRoute)
+}
+
+func stringsHasPortalPrefix(path string) bool {
+ return len(path) >= len(portalRoutePrefix) && path[:len(portalRoutePrefix)] == portalRoutePrefix
+}
diff --git a/core/server/web/query_flash.go b/core/server/web/query_flash.go
index 5be979b8..8e59a587 100644
--- a/core/server/web/query_flash.go
+++ b/core/server/web/query_flash.go
@@ -14,10 +14,6 @@ var importFlashPattern = regexp.MustCompile(`^imported_(\d+)_updated_(\d+)_skipp
// FlashFromQueryMessage converts ?msg= query values into workspace flash banners.
func FlashFromQueryMessage(msg string) (render.FlashMessage, bool) {
- return flashFromQueryMessage(msg)
-}
-
-func flashFromQueryMessage(msg string) (render.FlashMessage, bool) {
msg = strings.TrimSpace(msg)
if msg == "" {
return render.FlashMessage{}, false
@@ -43,6 +39,20 @@ func flashFromQueryMessage(msg string) (render.FlashMessage, bool) {
switch msg {
case resetPasswordMsg:
return render.FlashMessage{Kind: "info", Title: "Password reset", Body: "If the account exists, reset instructions were sent."}, true
+ case oauthDeniedMsg:
+ return render.FlashMessage{Kind: "error", Title: "Sign-in failed", Body: "Single sign-on could not complete. Try again or use your password if enabled."}, true
+ case authLocalDisabledMsg:
+ return render.FlashMessage{Kind: "error", Title: "Password sign-in disabled", Body: "Use one of the sign-in providers below."}, true
+ case "totp_enroll_started":
+ return render.FlashMessage{Kind: "info", Title: "Scan the QR code", Body: "Add the account in your authenticator app, then enter the 6-digit code to enable two-factor authentication."}, true
+ case "totp_enabled":
+ return render.FlashMessage{Kind: "success", Title: "Two-factor enabled", Body: "You will need an authenticator code when signing in on new devices."}, true
+ case "totp_disabled":
+ return render.FlashMessage{Kind: "success", Title: "Two-factor disabled", Body: "Authenticator codes are no longer required for your account."}, true
+ case "totp_invalid":
+ return render.FlashMessage{Kind: "error", Title: "Invalid code", Body: "Check the 6-digit code from your authenticator app and try again."}, true
+ case "totp_enroll_failed":
+ return render.FlashMessage{Kind: "error", Title: "Could not start enrollment", Body: "Try again or contact your administrator."}, true
case "password_updated":
return render.FlashMessage{Kind: "success", Title: "Password updated", Body: "Your password was changed."}, true
case "password_mismatch":
@@ -71,18 +81,42 @@ func flashFromQueryMessage(msg string) (render.FlashMessage, bool) {
return render.FlashMessage{Kind: "success", Title: "Updated", Body: "Stage updated.", ToastOnly: true}, true
default:
if strings.HasPrefix(msg, "error:") {
- body := strings.TrimPrefix(msg, "error:")
+ body := sanitizeFlashQueryBody(strings.TrimPrefix(msg, "error:"))
return render.FlashMessage{Kind: "error", Title: "Error", Body: body}, true
}
if strings.HasPrefix(msg, "save_error:") {
- body := strings.TrimPrefix(msg, "save_error:")
+ body := sanitizeFlashQueryBody(strings.TrimPrefix(msg, "save_error:"))
return render.FlashMessage{Kind: "error", Title: "Save failed", Body: body}, true
}
if strings.HasPrefix(msg, "installed_") || strings.HasPrefix(msg, "uninstalled_") || strings.HasPrefix(msg, "upgraded_") {
return render.FlashMessage{Kind: "success", Title: "Apps", Body: strings.ReplaceAll(msg, "_", " ")}, true
}
- return render.FlashMessage{Kind: "info", Title: "", Body: msg}, true
+ return render.FlashMessage{Kind: "info", Title: "", Body: sanitizeFlashQueryBody(msg)}, true
+ }
+}
+
+func sanitizeFlashQueryBody(s string) string {
+ s = strings.TrimSpace(s)
+ if s == "" {
+ return s
+ }
+ const maxRunes = 500
+ runes := []rune(s)
+ if len(runes) > maxRunes {
+ runes = runes[:maxRunes]
+ }
+ out := make([]rune, 0, len(runes))
+ for _, r := range runes {
+ if r == '\n' || r == '\r' {
+ out = append(out, ' ')
+ continue
+ }
+ if r < 0x20 {
+ continue
+ }
+ out = append(out, r)
}
+ return strings.TrimSpace(string(out))
}
func appendQueryFlashesToViewRecord(r *http.Request, viewRecord *render.ViewRecordData) {
@@ -93,7 +127,7 @@ func appendQueryFlashesToViewRecord(r *http.Request, viewRecord *render.ViewReco
if msg == "" {
return
}
- if flash, ok := flashFromQueryMessage(msg); ok {
+ if flash, ok := FlashFromQueryMessage(msg); ok {
viewRecord.FlashMessages = append(viewRecord.FlashMessages, flash)
}
}
diff --git a/core/server/web/rate_limit.go b/core/server/web/rate_limit.go
index 0cb476cb..18fc2118 100644
--- a/core/server/web/rate_limit.go
+++ b/core/server/web/rate_limit.go
@@ -1,6 +1,74 @@
package web
-import "time"
+import (
+ "net/http"
+ "strings"
+ "sync"
+ "time"
+
+ "sumeru/core/server/config"
+)
+
+type rateBucket struct {
+ count int
+ window time.Time
+}
+
+var (
+ rateMu sync.Mutex
+ rateByIP = map[string]*rateBucket{}
+ rateLimitOn bool
+ rateLimitRPM int
+)
+
+// InitRateLimit reads rate_limit_rpm from loaded config.
+func InitRateLimit() {
+ rateLimitRPM = config.AppConfig.RateLimitRPM
+ rateLimitOn = rateLimitRPM > 0
+}
+
+func rateLimitedPath(path string) bool {
+ switch path {
+ case apiRPCRoute, loginRoute, totpLoginRoute, oauthStartRoute, oauthCallbackRoute:
+ return true
+ default:
+ return false
+ }
+}
+
+func allowRate(clientIP string) bool {
+ if !rateLimitOn {
+ return true
+ }
+ clientIP = strings.TrimSpace(clientIP)
+ if clientIP == "" {
+ clientIP = "unknown"
+ }
+ now := time.Now()
+ rateMu.Lock()
+ defer rateMu.Unlock()
+ b, ok := rateByIP[clientIP]
+ if !ok || now.Sub(b.window) >= time.Minute {
+ rateByIP[clientIP] = &rateBucket{count: 1, window: now}
+ return true
+ }
+ if b.count >= rateLimitRPM {
+ return false
+ }
+ b.count++
+ return true
+}
+
+func enforceRateLimit(w http.ResponseWriter, r *http.Request) bool {
+ if !rateLimitedPath(r.URL.Path) {
+ return true
+ }
+ if allowRate(clientIP(r)) {
+ return true
+ }
+ http.Error(w, "Too Many Requests", http.StatusTooManyRequests)
+ return false
+}
// pruneAttemptsAfter keeps timestamps strictly after cutoff (in-place slice reuse).
func pruneAttemptsAfter(attempts []time.Time, cutoff time.Time) []time.Time {
diff --git a/core/server/web/ratelimit.go b/core/server/web/ratelimit.go
deleted file mode 100644
index d50c6dd9..00000000
--- a/core/server/web/ratelimit.go
+++ /dev/null
@@ -1,71 +0,0 @@
-package web
-
-import (
- "net/http"
- "strings"
- "sync"
- "time"
-
- "sumeru/core/server/config"
-)
-
-type rateBucket struct {
- count int
- window time.Time
-}
-
-var (
- rateMu sync.Mutex
- rateByIP = map[string]*rateBucket{}
- rateLimitOn bool
- rateLimitRPM int
-)
-
-// InitRateLimit reads rate_limit_rpm from loaded config.
-func InitRateLimit() {
- rateLimitRPM = config.AppConfig.RateLimitRPM
- rateLimitOn = rateLimitRPM > 0
-}
-
-func rateLimitedPath(path string) bool {
- switch path {
- case apiRPCRoute, loginRoute:
- return true
- default:
- return false
- }
-}
-
-func allowRate(clientIP string) bool {
- if !rateLimitOn {
- return true
- }
- clientIP = strings.TrimSpace(clientIP)
- if clientIP == "" {
- clientIP = "unknown"
- }
- now := time.Now()
- rateMu.Lock()
- defer rateMu.Unlock()
- b, ok := rateByIP[clientIP]
- if !ok || now.Sub(b.window) >= time.Minute {
- rateByIP[clientIP] = &rateBucket{count: 1, window: now}
- return true
- }
- if b.count >= rateLimitRPM {
- return false
- }
- b.count++
- return true
-}
-
-func enforceRateLimit(w http.ResponseWriter, r *http.Request) bool {
- if !rateLimitedPath(r.URL.Path) {
- return true
- }
- if allowRate(clientIP(r)) {
- return true
- }
- http.Error(w, "Too Many Requests", http.StatusTooManyRequests)
- return false
-}
diff --git a/core/server/web/routes_table.go b/core/server/web/routes_table.go
index 3cc59e0b..7ad6b138 100644
--- a/core/server/web/routes_table.go
+++ b/core/server/web/routes_table.go
@@ -58,6 +58,8 @@ func registerAuthRoutes() {
registerPublic(http.MethodGet, logoutRoute, LogoutGet)
registerSession(http.MethodPost, logoutRoute, LogoutPost)
registerAPIKeyRevealRoute()
+ registerOAuthRoutes()
+ registerTOTPRoutes()
}
func registerWorkspaceRoutes() {
diff --git a/core/server/web/settings_account.go b/core/server/web/settings_account.go
index 4093029e..2199eb59 100644
--- a/core/server/web/settings_account.go
+++ b/core/server/web/settings_account.go
@@ -14,6 +14,7 @@ const settingsAccountRoute = "/web/settings/account"
func registerSettingsAccountRoutes() {
registerSession(http.MethodGet, settingsAccountRoute, SettingsAccountGetHandler)
registerSession(http.MethodPost, settingsAccountRoute, SettingsAccountPostHandler)
+ registerSettingsAccountTOTPRoutes()
}
func SettingsAccountGetHandler(w http.ResponseWriter, r *http.Request) {
@@ -28,8 +29,13 @@ func SettingsAccountGetHandler(w http.ResponseWriter, r *http.Request) {
if !ok {
return
}
- flash, _ := flashFromQueryMessage(r.URL.Query().Get("msg"))
- renderSettingsAccountPage(w, r, settingsAccountData{CSRFToken: CSRFTokenForRequest(r), Flash: flash}, menuIDStr)
+ flash, _ := FlashFromQueryMessage(r.URL.Query().Get("msg"))
+ actor := orm.SecurityUID(ctx)
+ renderSettingsAccountPage(w, r, settingsAccountData{
+ CSRFToken: CSRFTokenForRequest(r),
+ Flash: flash,
+ TOTP: loadSettingsAccountTOTPState(ctx, actor),
+ }, menuIDStr)
}
func SettingsAccountPostHandler(w http.ResponseWriter, r *http.Request) {
@@ -49,6 +55,10 @@ func SettingsAccountPostHandler(w http.ResponseWriter, r *http.Request) {
}
ctx := r.Context()
actor := orm.SecurityUID(ctx)
+ action := strings.TrimSpace(r.PostForm.Get("action"))
+ if handleSettingsAccountTOTP(w, r, ctx, actor, action) {
+ return
+ }
pw := strings.TrimSpace(r.PostForm.Get("password_plain"))
confirm := strings.TrimSpace(r.PostForm.Get("password_plain_confirm"))
if pw == "" {
@@ -69,6 +79,7 @@ func SettingsAccountPostHandler(w http.ResponseWriter, r *http.Request) {
type settingsAccountData struct {
CSRFToken string
Flash render.FlashMessage
+ TOTP settingsAccountTOTPState
}
func renderSettingsAccountPage(w http.ResponseWriter, r *http.Request, pageData settingsAccountData, menuIDStr string) {
diff --git a/core/server/web/settings_account_totp.go b/core/server/web/settings_account_totp.go
new file mode 100644
index 00000000..ae804dda
--- /dev/null
+++ b/core/server/web/settings_account_totp.go
@@ -0,0 +1,104 @@
+package web
+
+import (
+ "context"
+ "net/http"
+ "strings"
+
+ "sumeru/core/orm"
+ "sumeru/core/server/auth"
+)
+
+const settingsAccountTOTPQRRoute = "/web/settings/account/totp-qr"
+
+func registerSettingsAccountTOTPRoutes() {
+ registerSession(http.MethodGet, settingsAccountTOTPQRRoute, SettingsAccountTOTPQrHandler)
+}
+
+func SettingsAccountTOTPQrHandler(w http.ResponseWriter, r *http.Request) {
+ if !requireLogin(w, r) {
+ return
+ }
+ ctx := r.Context()
+ actor := orm.SecurityUID(ctx)
+ secret, ok, err := orm.OwnTOTPEnrollmentPending(ctx, actor, actor)
+ if err != nil || !ok {
+ http.Error(w, "not found", http.StatusNotFound)
+ return
+ }
+ login := userLoginName(ctx, actor)
+ uri := auth.TOTPOtpauthURI("Sumeru", login, secret)
+ png, err := auth.TOTPQRPng(uri, 220)
+ if err != nil {
+ http.Error(w, "qr unavailable", http.StatusInternalServerError)
+ return
+ }
+ w.Header().Set("Content-Type", "image/png")
+ w.Header().Set("Cache-Control", "no-store")
+ _, _ = w.Write(png)
+}
+
+func userLoginName(ctx context.Context, userID int) string {
+ rec, err := orm.SearchOne(ctx, "core.user", map[string]interface{}{"id": userID})
+ if err != nil || rec == nil {
+ return "user"
+ }
+ login := strings.TrimSpace(ormRowString(rec, "login"))
+ if login == "" {
+ return "user"
+ }
+ return login
+}
+
+func ormRowString(rec map[string]interface{}, key string) string {
+ if v, ok := rec[key].(string); ok {
+ return v
+ }
+ return ""
+}
+
+func handleSettingsAccountTOTP(w http.ResponseWriter, r *http.Request, ctx context.Context, actor int, action string) bool {
+ switch strings.TrimSpace(action) {
+ case "totp_enroll_start":
+ if _, err := orm.BeginOwnTOTPEnrollment(ctx, actor, actor); err != nil {
+ redirectWithWebMessage(w, r, settingsAccountRoute, "totp_enroll_failed")
+ return true
+ }
+ redirectWithWebMessage(w, r, settingsAccountRoute, "totp_enroll_started")
+ return true
+ case "totp_enroll_confirm":
+ code := strings.TrimSpace(r.PostForm.Get("totp_code"))
+ if err := orm.ConfirmOwnTOTPEnrollment(ctx, actor, actor, code); err != nil {
+ redirectWithWebMessage(w, r, settingsAccountRoute, "totp_invalid")
+ return true
+ }
+ redirectWithWebMessage(w, r, settingsAccountRoute, "totp_enabled")
+ return true
+ case "totp_disable":
+ code := strings.TrimSpace(r.PostForm.Get("totp_code"))
+ if err := orm.DisableOwnTOTP(ctx, actor, actor, code); err != nil {
+ redirectWithWebMessage(w, r, settingsAccountRoute, "totp_invalid")
+ return true
+ }
+ redirectWithWebMessage(w, r, settingsAccountRoute, "totp_disabled")
+ return true
+ default:
+ return false
+ }
+}
+
+func loadSettingsAccountTOTPState(ctx context.Context, actor int) settingsAccountTOTPState {
+ enabled, _ := orm.UserTOTPEnabled(ctx, actor)
+ pending := false
+ if !enabled {
+ if _, ok, _ := orm.OwnTOTPEnrollmentPending(ctx, actor, actor); ok {
+ pending = true
+ }
+ }
+ return settingsAccountTOTPState{Enabled: enabled, EnrollPending: pending}
+}
+
+type settingsAccountTOTPState struct {
+ Enabled bool
+ EnrollPending bool
+}
diff --git a/core/server/web/settings_field_acl.go b/core/server/web/settings_field_acl.go
index c5be0507..36cb5038 100644
--- a/core/server/web/settings_field_acl.go
+++ b/core/server/web/settings_field_acl.go
@@ -66,7 +66,7 @@ func SettingsFieldACLGetHandler(w http.ResponseWriter, r *http.Request) {
menuID := resolveSettingsMenuXMLID(ctx, render.MenuFieldAccessMatrixXMLID, rootMenuID)
q := r.URL.Query()
model := strings.TrimSpace(q.Get("model"))
- flash, _ := flashFromQueryMessage(q.Get("msg"))
+ flash, _ := FlashFromQueryMessage(q.Get("msg"))
page, err := buildFieldACLMatrixPage(ctx, model, q.Get("group_q"), q.Get("field_q"), q.Get("show_all") == "1")
if err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
diff --git a/core/server/web/settings_model_acl.go b/core/server/web/settings_model_acl.go
index e9dbe534..a468dce1 100644
--- a/core/server/web/settings_model_acl.go
+++ b/core/server/web/settings_model_acl.go
@@ -54,7 +54,7 @@ func SettingsModelACLGetHandler(w http.ResponseWriter, r *http.Request) {
menuID := resolveSettingsMenuXMLID(ctx, render.MenuModelAccessMatrixXMLID, rootMenuID)
q := r.URL.Query()
model := strings.TrimSpace(q.Get("model"))
- flash, _ := flashFromQueryMessage(q.Get("msg"))
+ flash, _ := FlashFromQueryMessage(q.Get("msg"))
page, err := buildModelACLMatrixPage(ctx, model, q.Get("group_q"), q.Get("show_all") == "1")
if err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
diff --git a/core/server/web/signed_cookie.go b/core/server/web/signed_cookie.go
new file mode 100644
index 00000000..91bc24f3
--- /dev/null
+++ b/core/server/web/signed_cookie.go
@@ -0,0 +1,73 @@
+package web
+
+import (
+ "crypto/hmac"
+ "crypto/sha256"
+ "encoding/base64"
+ "encoding/json"
+ "strings"
+ "time"
+
+ "sumeru/core/server/config"
+)
+
+func webMACSecret() []byte {
+ secret := strings.TrimSpace(config.AppConfig.CSRFSecret)
+ if secret == "" {
+ secret = "dev-web-cookie-mac"
+ }
+ return []byte(secret)
+}
+
+func signWebCookiePayload(payload []byte) string {
+ mac := hmac.New(sha256.New, webMACSecret())
+ _, _ = mac.Write(payload)
+ sig := base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
+ body := base64.RawURLEncoding.EncodeToString(payload)
+ return body + "." + sig
+}
+
+func verifyWebCookiePayload(value string, out interface{}) bool {
+ parts := strings.Split(strings.TrimSpace(value), ".")
+ if len(parts) != 2 {
+ return false
+ }
+ payload, err := base64.RawURLEncoding.DecodeString(parts[0])
+ if err != nil {
+ return false
+ }
+ sig, err := base64.RawURLEncoding.DecodeString(parts[1])
+ if err != nil {
+ return false
+ }
+ mac := hmac.New(sha256.New, webMACSecret())
+ _, _ = mac.Write(payload)
+ if !hmac.Equal(sig, mac.Sum(nil)) {
+ return false
+ }
+ return json.Unmarshal(payload, out) == nil
+}
+
+type signedUIDPayload struct {
+ UID int `json:"uid"`
+ Exp int64 `json:"exp"`
+}
+
+func mintSignedUIDCookieValue(uid int, ttl time.Duration) string {
+ payload, _ := json.Marshal(signedUIDPayload{
+ UID: uid,
+ Exp: time.Now().Add(ttl).Unix(),
+ })
+ return signWebCookiePayload(payload)
+}
+
+func parseSignedUIDCookie(value string) (uid int, ok bool) {
+ var p signedUIDPayload
+ if !verifyWebCookiePayload(value, &p) {
+ return 0, false
+ }
+ if p.UID <= 0 || p.Exp <= 0 || time.Now().Unix() > p.Exp {
+ return 0, false
+ }
+ return p.UID, true
+}
diff --git a/core/server/web/swc_bus.go b/core/server/web/swc_bus.go
index 99367bab..aa773ea7 100644
--- a/core/server/web/swc_bus.go
+++ b/core/server/web/swc_bus.go
@@ -1,13 +1,44 @@
package web
import (
+ "context"
+ "encoding/json"
+ "net"
"net/http"
+ "net/url"
+ "strings"
+ "sync"
+ "time"
"github.com/gorilla/websocket"
+ "github.com/lib/pq"
+ "sumeru/core/applog"
+ "sumeru/core/orm"
+ "sumeru/core/queue"
)
const swcBusRoute = "/web/swc/bus"
+var (
+ swcBusUpgrader = websocket.Upgrader{
+ CheckOrigin: checkSwcBusOrigin,
+ }
+ globalBusHub *busHub
+ globalBusHubOnce sync.Once
+)
+
+const (
+ swcBusSendBuffer = 32
+ swcBusPingInterval = 45 * time.Second
+ swcBusPongWait = 60 * time.Second
+)
+
+func init() {
+ orm.SetBusPublishHook(func(ctx context.Context, channel string, payload map[string]interface{}) {
+ _ = PublishSwcBusEvent(ctx, channel, payload)
+ })
+}
+
func registerSwcBusRoute() {
registerSession(http.MethodGet, swcBusRoute, SwcBusHandler)
}
@@ -23,3 +54,297 @@ func SwcBusHandler(w http.ResponseWriter, r *http.Request) {
}
serveSwcBusWebSocket(w, r, AuthenticatedUserID(r))
}
+
+// StartBusListener listens for PostgreSQL NOTIFY and dispatches bus events locally.
+func StartBusListener(parent context.Context, connStr string) {
+ if parent == nil || connStr == "" || orm.DB == nil {
+ return
+ }
+ go func() {
+ backoff := time.Second
+ for {
+ if parent.Err() != nil {
+ return
+ }
+ if err := listenBusLoop(parent, connStr); err != nil {
+ applog.WarnMsg(parent, "web", "bus_listen", "bus listener ended", err, nil)
+ time.Sleep(backoff)
+ if backoff < 30*time.Second {
+ backoff *= 2
+ }
+ continue
+ }
+ return
+ }
+ }()
+}
+
+func listenBusLoop(parent context.Context, connStr string) error {
+ listener := pq.NewListener(connStr, 10*time.Second, time.Minute, func(ev pq.ListenerEventType, err error) {
+ if err != nil {
+ applog.WarnMsg(parent, "web", "bus_listen", "pq listener event", err, map[string]interface{}{"event": int(ev)})
+ }
+ })
+ defer func() { _ = listener.Close() }()
+ if err := listener.Listen(orm.BusNotifyChannel()); err != nil {
+ return err
+ }
+ for {
+ select {
+ case <-parent.Done():
+ return parent.Err()
+ case n, ok := <-listener.Notify:
+ if !ok {
+ return nil
+ }
+ if n != nil {
+ dispatchBusEventByID(parent, parseBusEventIDNotify(n.Extra))
+ }
+ }
+ }
+}
+
+func checkSwcBusOrigin(r *http.Request) bool {
+ origin := strings.TrimSpace(r.Header.Get("Origin"))
+ if origin == "" {
+ return true
+ }
+ u, err := url.Parse(origin)
+ if err != nil {
+ return false
+ }
+ reqHost := r.Host
+ if h, _, err := net.SplitHostPort(reqHost); err == nil {
+ reqHost = h
+ }
+ return strings.EqualFold(u.Hostname(), reqHost)
+}
+
+type swcBusClient struct {
+ uid int
+ conn *websocket.Conn
+ send chan []byte
+ channels map[string]struct{}
+ subMu sync.Mutex
+}
+
+type busHub struct {
+ mu sync.RWMutex
+ clients map[*swcBusClient]struct{}
+}
+
+func ensureBusHub() *busHub {
+ globalBusHubOnce.Do(func() {
+ globalBusHub = &busHub{clients: make(map[*swcBusClient]struct{})}
+ queue.Subscribe("outbox", handleOutboxBusBridge)
+ })
+ return globalBusHub
+}
+
+func handleOutboxBusBridge(ctx context.Context, msg queue.Message) error {
+ var envelope map[string]interface{}
+ if err := json.Unmarshal(msg.Payload, &envelope); err != nil {
+ return nil
+ }
+ name, _ := envelope["name"].(string)
+ if name == "" {
+ return nil
+ }
+ inner, _ := envelope["payload"].(map[string]interface{})
+ if inner == nil {
+ inner = map[string]interface{}{}
+ }
+ model, _ := inner["model"].(string)
+ var rid int
+ switch v := inner["id"].(type) {
+ case float64:
+ rid = int(v)
+ case int:
+ rid = v
+ }
+ if model != "" && rid > 0 && strings.HasPrefix(name, "record.") {
+ PublishRecordBusEvent(ctx, name, model, rid)
+ }
+ return nil
+}
+
+func (h *busHub) register(c *swcBusClient) {
+ h.mu.Lock()
+ h.clients[c] = struct{}{}
+ h.mu.Unlock()
+}
+
+func (h *busHub) unregister(c *swcBusClient) {
+ h.mu.Lock()
+ delete(h.clients, c)
+ h.mu.Unlock()
+}
+
+func (h *busHub) publishChannel(channel string, msg []byte) {
+ h.mu.RLock()
+ targets := make([]*swcBusClient, 0)
+ for c := range h.clients {
+ c.subMu.Lock()
+ _, ok := c.channels[channel]
+ c.subMu.Unlock()
+ if ok {
+ targets = append(targets, c)
+ }
+ }
+ h.mu.RUnlock()
+ for _, c := range targets {
+ select {
+ case c.send <- msg:
+ default:
+ }
+ }
+}
+
+func (c *swcBusClient) subscribeChannels(ctx context.Context, channels []string, afterID int64) {
+ for _, ch := range channels {
+ ch = strings.TrimSpace(ch)
+ if ch == "" || !AuthorizeSwcBusChannel(ctx, c.uid, ch) {
+ continue
+ }
+ c.subMu.Lock()
+ if len(c.channels) >= maxBusSubscriptions {
+ c.subMu.Unlock()
+ break
+ }
+ c.channels[ch] = struct{}{}
+ c.subMu.Unlock()
+ if afterID > 0 {
+ rows, err := orm.ListBusEventsAfter(ctx, []string{ch}, afterID, 500)
+ if err != nil {
+ continue
+ }
+ for _, row := range rows {
+ frame, err := marshalBusEventFrame(row.ID, row.Channel, row.Payload)
+ if err != nil {
+ continue
+ }
+ select {
+ case c.send <- frame:
+ default:
+ }
+ }
+ }
+ }
+}
+
+func (c *swcBusClient) unsubscribeChannels(channels []string) {
+ c.subMu.Lock()
+ defer c.subMu.Unlock()
+ for _, ch := range channels {
+ ch = strings.TrimSpace(ch)
+ delete(c.channels, ch)
+ }
+}
+
+func (c *swcBusClient) writePump() {
+ ticker := time.NewTicker(swcBusPingInterval)
+ defer func() {
+ ticker.Stop()
+ c.conn.Close()
+ }()
+ for {
+ select {
+ case msg, ok := <-c.send:
+ if !ok {
+ return
+ }
+ _ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
+ if err := c.conn.WriteMessage(websocket.TextMessage, msg); err != nil {
+ return
+ }
+ case <-ticker.C:
+ _ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
+ if err := c.conn.WriteMessage(websocket.PingMessage, nil); err != nil {
+ return
+ }
+ }
+ }
+}
+
+func (c *swcBusClient) readPump(h *busHub, baseCtx context.Context) {
+ defer func() {
+ h.unregister(c)
+ close(c.send)
+ c.conn.Close()
+ }()
+ _ = c.conn.SetReadDeadline(time.Now().Add(swcBusPongWait))
+ c.conn.SetPongHandler(func(string) error {
+ _ = c.conn.SetReadDeadline(time.Now().Add(swcBusPongWait))
+ return nil
+ })
+ for {
+ _, data, err := c.conn.ReadMessage()
+ if err != nil {
+ return
+ }
+ c.handleClientFrame(baseCtx, data)
+ }
+}
+
+func (c *swcBusClient) handleClientFrame(ctx context.Context, data []byte) {
+ var frame map[string]interface{}
+ if err := json.Unmarshal(data, &frame); err != nil {
+ return
+ }
+ typ, _ := frame["type"].(string)
+ switch typ {
+ case "subscribe":
+ channels := stringSliceField(frame["channels"])
+ afterID, _ := orm.CoerceInt64(frame["last_event_id"])
+ c.subscribeChannels(ctx, channels, afterID)
+ case "unsubscribe":
+ c.unsubscribeChannels(stringSliceField(frame["channels"]))
+ case "ping":
+ if b, err := marshalBusFrame(map[string]interface{}{"type": "pong"}); err == nil {
+ select {
+ case c.send <- b:
+ default:
+ }
+ }
+ }
+}
+
+func stringSliceField(v interface{}) []string {
+ switch t := v.(type) {
+ case []interface{}:
+ out := make([]string, 0, len(t))
+ for _, item := range t {
+ if s, ok := item.(string); ok && s != "" {
+ out = append(out, s)
+ }
+ }
+ return out
+ case []string:
+ return t
+ default:
+ return nil
+ }
+}
+
+func marshalBusFrame(m map[string]interface{}) ([]byte, error) {
+ return json.Marshal(m)
+}
+
+func serveSwcBusWebSocket(w http.ResponseWriter, r *http.Request, uid int) {
+ hub := ensureBusHub()
+ conn, err := swcBusUpgrader.Upgrade(w, r, nil)
+ if err != nil {
+ return
+ }
+ client := &swcBusClient{
+ uid: uid,
+ conn: conn,
+ send: make(chan []byte, swcBusSendBuffer),
+ channels: make(map[string]struct{}),
+ }
+ hub.register(client)
+ ctx := r.Context()
+ client.subscribeChannels(ctx, []string{UserNotificationsChannel(uid)}, 0)
+ go client.writePump()
+ client.readPump(hub, ctx)
+}
diff --git a/core/server/web/swc_bus_channel.go b/core/server/web/swc_bus_channel.go
new file mode 100644
index 00000000..cc724079
--- /dev/null
+++ b/core/server/web/swc_bus_channel.go
@@ -0,0 +1,79 @@
+package web
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+ "strings"
+
+ "sumeru/core/orm"
+)
+
+const maxBusSubscriptions = 64
+
+// AuthorizeSwcBusChannel checks whether uid may subscribe to channel.
+func AuthorizeSwcBusChannel(ctx context.Context, uid int, channel string) bool {
+ if uid <= 0 || channel == "" {
+ return false
+ }
+ channel = strings.TrimSpace(channel)
+ parts := strings.Split(channel, "/")
+ if len(parts) < 2 {
+ return false
+ }
+ switch parts[0] {
+ case "user":
+ id, err := strconv.Atoi(parts[1])
+ return err == nil && id == uid
+ case "group":
+ if len(parts) < 2 {
+ return false
+ }
+ xmlid := strings.Join(parts[1:], ".")
+ return orm.UserHasGroupXML(ctx, uid, xmlid)
+ case "record":
+ if len(parts) != 3 {
+ return false
+ }
+ model := parts[1]
+ rid, err := strconv.Atoi(parts[2])
+ if err != nil || rid <= 0 {
+ return false
+ }
+ if err := orm.CheckModelAccess(ctx, uid, model, "read"); err != nil {
+ return false
+ }
+ _, err = orm.SearchOne(ctx, model, map[string]interface{}{"id": rid})
+ return err == nil
+ case "model":
+ if len(parts) != 2 {
+ return false
+ }
+ return orm.CheckModelAccess(ctx, uid, parts[1], "read") == nil
+ case "company":
+ if len(parts) != 2 {
+ return false
+ }
+ cid, err := strconv.Atoi(parts[1])
+ if err != nil || cid <= 0 {
+ return false
+ }
+ return orm.UserAllowedCompany(ctx, uid, int64(cid))
+ default:
+ return false
+ }
+}
+
+// RecordBusChannel builds the standard record mutation channel name.
+func RecordBusChannel(model string, id int) string {
+ model = strings.TrimSpace(model)
+ if model == "" || id <= 0 {
+ return ""
+ }
+ return fmt.Sprintf("record/%s/%d", model, id)
+}
+
+// UserNotificationsChannel is the inbox channel for uid.
+func UserNotificationsChannel(uid int) string {
+ return orm.UserNotificationsBusChannel(uid)
+}
diff --git a/core/server/web/swc_bus_hub.go b/core/server/web/swc_bus_hub.go
deleted file mode 100644
index 7bba43c8..00000000
--- a/core/server/web/swc_bus_hub.go
+++ /dev/null
@@ -1,149 +0,0 @@
-package web
-
-import (
- "context"
- "encoding/json"
- "net"
- "net/http"
- "net/url"
- "strings"
- "sync"
-
- "github.com/gorilla/websocket"
- "sumeru/core/orm"
- "sumeru/core/queue"
-)
-
-var (
- swcBusUpgrader = websocket.Upgrader{
- CheckOrigin: checkSwcBusOrigin,
- }
- globalBusHub *busHub
- globalBusHubOnce sync.Once
-)
-
-func checkSwcBusOrigin(r *http.Request) bool {
- origin := strings.TrimSpace(r.Header.Get("Origin"))
- if origin == "" {
- return true
- }
- u, err := url.Parse(origin)
- if err != nil {
- return false
- }
- reqHost := r.Host
- if h, _, err := net.SplitHostPort(reqHost); err == nil {
- reqHost = h
- }
- originHost := u.Hostname()
- return strings.EqualFold(originHost, reqHost)
-}
-
-type swcBusClient struct {
- uid int
- conn *websocket.Conn
- send chan []byte
-}
-
-type busHub struct {
- mu sync.RWMutex
- clients map[*swcBusClient]struct{}
-}
-
-func ensureBusHub() *busHub {
- globalBusHubOnce.Do(func() {
- globalBusHub = &busHub{clients: make(map[*swcBusClient]struct{})}
- queue.Subscribe("outbox", func(ctx context.Context, msg queue.Message) error {
- var envelope map[string]interface{}
- if err := json.Unmarshal(msg.Payload, &envelope); err != nil {
- return nil
- }
- name, _ := envelope["name"].(string)
- if name == "" {
- return nil
- }
- inner, _ := envelope["payload"].(map[string]interface{})
- if inner == nil {
- inner = map[string]interface{}{}
- }
- actor, _ := orm.CoerceInt64(envelope["actor"])
- out, err := json.Marshal(map[string]interface{}{
- "channel": name,
- "payload": inner,
- })
- if err != nil {
- return nil
- }
- globalBusHub.broadcast(int(actor), out)
- return nil
- })
- })
- return globalBusHub
-}
-
-func (h *busHub) register(c *swcBusClient) {
- h.mu.Lock()
- h.clients[c] = struct{}{}
- h.mu.Unlock()
-}
-
-func (h *busHub) unregister(c *swcBusClient) {
- h.mu.Lock()
- delete(h.clients, c)
- h.mu.Unlock()
-}
-
-func (h *busHub) broadcast(actor int, msg []byte) {
- h.mu.RLock()
- targets := make([]*swcBusClient, 0, len(h.clients))
- for c := range h.clients {
- if actor <= 0 || c.uid == actor {
- targets = append(targets, c)
- }
- }
- h.mu.RUnlock()
- for _, c := range targets {
- select {
- case c.send <- msg:
- default:
- }
- }
-}
-
-func (c *swcBusClient) writePump() {
- defer c.conn.Close()
- for msg := range c.send {
- if err := c.conn.WriteMessage(websocket.TextMessage, msg); err != nil {
- return
- }
- }
-}
-
-func (c *swcBusClient) readPump(h *busHub) {
- defer func() {
- h.unregister(c)
- close(c.send)
- c.conn.Close()
- }()
- for {
- if _, _, err := c.conn.ReadMessage(); err != nil {
- return
- }
- }
-}
-
-func serveSwcBusWebSocket(w http.ResponseWriter, r *http.Request, uid int) {
- hub := ensureBusHub()
- conn, err := swcBusUpgrader.Upgrade(w, r, nil)
- if err != nil {
- return
- }
- client := &swcBusClient{
- uid: uid,
- conn: conn,
- send: make(chan []byte, 16),
- }
- hub.register(client)
- go client.writePump()
- client.readPump(hub)
-}
diff --git a/core/server/web/swc_bus_publish.go b/core/server/web/swc_bus_publish.go
new file mode 100644
index 00000000..18bfa5c4
--- /dev/null
+++ b/core/server/web/swc_bus_publish.go
@@ -0,0 +1,78 @@
+package web
+
+import (
+ "context"
+ "fmt"
+ "strconv"
+
+ "sumeru/core/orm"
+)
+
+// PublishSwcBusEvent persists and fan-outs a bus event to subscribers (all processes via NOTIFY).
+func PublishSwcBusEvent(ctx context.Context, channel string, payload map[string]interface{}) error {
+ if channel == "" {
+ return fmt.Errorf("bus channel required")
+ }
+ eventID, err := orm.InsertBusEvent(ctx, channel, payload)
+ if err != nil {
+ return err
+ }
+ if eventID <= 0 {
+ return nil
+ }
+ if err := orm.NotifyBusEvent(eventID); err != nil {
+ return err
+ }
+ dispatchLocalBusEvent(eventID, channel, payload)
+ return nil
+}
+
+// PublishRecordBusEvent emits record.created|updated|deleted to record and model channels.
+func PublishRecordBusEvent(ctx context.Context, eventName, model string, id int) {
+ if model == "" || id <= 0 {
+ return
+ }
+ payload := map[string]interface{}{
+ "event": eventName,
+ "model": model,
+ "id": id,
+ }
+ ch := RecordBusChannel(model, id)
+ _ = PublishSwcBusEvent(ctx, ch, payload)
+ modelCh := fmt.Sprintf("model/%s", model)
+ _ = PublishSwcBusEvent(ctx, modelCh, payload)
+}
+
+func dispatchLocalBusEvent(eventID int64, channel string, payload map[string]interface{}) {
+ hub := ensureBusHub()
+ frame, err := marshalBusEventFrame(eventID, channel, payload)
+ if err != nil {
+ return
+ }
+ hub.publishChannel(channel, frame)
+}
+
+func dispatchBusEventByID(ctx context.Context, eventID int64) {
+ if eventID <= 0 {
+ return
+ }
+ channel, payload, err := orm.LoadBusEventByID(ctx, eventID)
+ if err != nil || channel == "" {
+ return
+ }
+ dispatchLocalBusEvent(eventID, channel, payload)
+}
+
+func marshalBusEventFrame(eventID int64, channel string, payload map[string]interface{}) ([]byte, error) {
+ return marshalBusFrame(map[string]interface{}{
+ "type": "event",
+ "id": eventID,
+ "channel": channel,
+ "payload": payload,
+ })
+}
+
+func parseBusEventIDNotify(payload string) int64 {
+ id, _ := strconv.ParseInt(payload, 10, 64)
+ return id
+}
diff --git a/core/server/web/swc_notifications.go b/core/server/web/swc_notifications.go
new file mode 100644
index 00000000..9fdd807e
--- /dev/null
+++ b/core/server/web/swc_notifications.go
@@ -0,0 +1,150 @@
+package web
+
+import (
+ "net/http"
+ "strings"
+ "time"
+
+ "sumeru/core/orm"
+)
+
+const (
+ swcNotificationsRoute = "/web/swc/notifications"
+ swcNotificationsReadRoute = "/web/swc/notifications/read"
+)
+
+func registerSwcNotificationRoutes() {
+ registerSession(http.MethodGet, swcNotificationsRoute, SwcNotificationsHandler)
+ registerSession(http.MethodPost, swcNotificationsReadRoute, SwcNotificationsMarkReadHandler)
+}
+
+type swcNotificationItem struct {
+ ID int `json:"id"`
+ MessageID int `json:"messageId"`
+ Body string `json:"body"`
+ Author string `json:"author"`
+ RecordModel string `json:"recordModel"`
+ RecordID int `json:"recordId"`
+ IsRead bool `json:"isRead"`
+ CreateDate string `json:"createDate"`
+}
+
+// SwcNotificationsHandler GET /web/swc/notifications
+func SwcNotificationsHandler(w http.ResponseWriter, r *http.Request) {
+ if !requireLogin(w, r) {
+ return
+ }
+ uid := AuthenticatedUserID(r)
+ ctx := r.Context()
+ if err := orm.CheckModelAccess(ctx, uid, "mail.notification", "read"); err != nil {
+ http.Error(w, "forbidden", http.StatusForbidden)
+ return
+ }
+ unreadOnly := strings.TrimSpace(r.URL.Query().Get("unread")) == "1"
+ domain := [][]interface{}{{"user_id", "=", uid}}
+ if unreadOnly {
+ domain = append(domain, []interface{}{"is_read", "=", false})
+ }
+ rows, err := orm.SearchPage(ctx, "mail.notification", domain, 50, 0, "create_date DESC")
+ if err != nil {
+ http.Error(w, "load failed", http.StatusInternalServerError)
+ return
+ }
+ items := make([]swcNotificationItem, 0, len(rows))
+ unread := 0
+ for _, row := range rows {
+ msgID := int(coerceFloat(row["message_id"]))
+ body, author := "", ""
+ if msgID > 0 {
+ if msg, err := orm.SearchOne(ctx, "mail.message", map[string]interface{}{"id": msgID}); err == nil && msg != nil {
+ body, _ = msg["body"].(string)
+ author, _ = msg["author"].(string)
+ }
+ }
+ isRead, _ := row["is_read"].(bool)
+ if !isRead {
+ unread++
+ }
+ items = append(items, swcNotificationItem{
+ ID: int(coerceFloat(row["id"])),
+ MessageID: msgID,
+ Body: body,
+ Author: author,
+ RecordModel: notifStringField(row["record_model"]),
+ RecordID: int(coerceFloat(row["record_id"])),
+ IsRead: isRead,
+ CreateDate: notifStringField(row["create_date"]),
+ })
+ }
+ writeJSONResponse(w, map[string]interface{}{
+ "items": items,
+ "unread": unread,
+ })
+}
+
+// SwcNotificationsMarkReadHandler POST /web/swc/notifications/read
+func SwcNotificationsMarkReadHandler(w http.ResponseWriter, r *http.Request) {
+ if !requireLogin(w, r) {
+ return
+ }
+ if !ParsePostForm(w, r) {
+ return
+ }
+ uid := AuthenticatedUserID(r)
+ ctx := r.Context()
+ if err := orm.CheckModelAccess(ctx, uid, "mail.notification", "write"); err != nil {
+ http.Error(w, "forbidden", http.StatusForbidden)
+ return
+ }
+ markAll := strings.TrimSpace(r.PostFormValue("all")) == "1"
+ now := time.Now().UTC().Format(time.RFC3339)
+ vals := map[string]interface{}{"is_read": true, "read_date": now}
+ if markAll {
+ rows, _ := orm.Search(ctx, "mail.notification", [][]interface{}{
+ {"user_id", "=", uid},
+ {"is_read", "=", false},
+ })
+ for _, row := range rows {
+ id := int(coerceFloat(row["id"]))
+ _, _ = orm.Update(ctx, "mail.notification", [][]interface{}{{"id", "=", id}}, vals)
+ }
+ writeJSONResponse(w, map[string]interface{}{"ok": true})
+ return
+ }
+ id := int(coerceFloat(r.PostFormValue("id")))
+ if id <= 0 {
+ http.Error(w, "id required", http.StatusBadRequest)
+ return
+ }
+ rec, err := orm.SearchOne(ctx, "mail.notification", map[string]interface{}{"id": id})
+ if err != nil || rec == nil || int(coerceFloat(rec["user_id"])) != uid {
+ http.Error(w, "not found", http.StatusNotFound)
+ return
+ }
+ _, err = orm.Update(ctx, "mail.notification", [][]interface{}{{"id", "=", id}}, vals)
+ if err != nil {
+ http.Error(w, "update failed", http.StatusInternalServerError)
+ return
+ }
+ writeJSONResponse(w, map[string]interface{}{"ok": true})
+}
+
+func coerceFloat(v interface{}) float64 {
+ switch t := v.(type) {
+ case float64:
+ return t
+ case int:
+ return float64(t)
+ case int64:
+ return float64(t)
+ default:
+ return 0
+ }
+}
+
+func notifStringField(v interface{}) string {
+ if s, ok := v.(string); ok {
+ return s
+ }
+ return ""
+}
diff --git a/core/server/web/swc_workspace.go b/core/server/web/swc_workspace.go
index a133ac26..5cda26e6 100644
--- a/core/server/web/swc_workspace.go
+++ b/core/server/web/swc_workspace.go
@@ -12,6 +12,7 @@ func registerSwcRoutes() {
registerSwcSavedSearchRoutes()
registerSwcBusRoute()
registerSwcChatterRoute()
+ registerSwcNotificationRoutes()
}
// SwcWorkspaceHandler GET /web/swc/workspace — JSON workspace payload for SWC.
diff --git a/core/server/web/testexports.go b/core/server/web/testexports.go
index 08f347b2..fe8253f3 100644
--- a/core/server/web/testexports.go
+++ b/core/server/web/testexports.go
@@ -52,6 +52,8 @@ const (
TestAppsCategoryField = appsCategoryField
TestAppsGroupByField = appsGroupByField
TestFlashMessageParam = flashMessageParam
+ TestOAuthDeniedMsg = oauthDeniedMsg
+ TestAuthLocalDisabledMsg = authLocalDisabledMsg
TestResetPasswordMsg = resetPasswordMsg
TestResetPasswordRoute = resetPasswordRoute
TestRecordModelField = recordModelField
@@ -418,6 +420,31 @@ func SessionCookieFromRequestForTest(r *http.Request) (sid, cookieName string) {
// LoginGetForTest exposes the login page GET handler for tests.
func LoginGetForTest(w http.ResponseWriter, r *http.Request) { LoginGet(w, r) }
+// OAuthCallbackForTest exposes the OAuth callback handler for tests.
+func OAuthCallbackForTest(w http.ResponseWriter, r *http.Request) { OAuthCallbackHandler(w, r) }
+
+// OAuthStartForTest exposes the OAuth start handler for tests.
+func OAuthStartForTest(w http.ResponseWriter, r *http.Request) { OAuthStartHandler(w, r) }
+
+// AuthLocalEnabledForTest reports whether password login is shown on the login page.
+func AuthLocalEnabledForTest(ctx context.Context) bool { return authLocalEnabled(ctx) }
+
+// LoginAuthProviderForTest is a login SSO button descriptor for tests.
+type LoginAuthProviderForTest loginAuthProvider
+
+// ListEnabledAuthProvidersForTest returns SSO providers for the login page.
+func ListEnabledAuthProvidersForTest(ctx context.Context) []LoginAuthProviderForTest {
+ rows := listEnabledAuthProviders(ctx)
+ out := make([]LoginAuthProviderForTest, len(rows))
+ for i, row := range rows {
+ out[i] = LoginAuthProviderForTest(row)
+ }
+ return out
+}
+
+// LoginPageCompanyNameForTest returns branding text for the login page.
+func LoginPageCompanyNameForTest(ctx context.Context) string { return loginPageCompanyName(ctx) }
+
// LogoutGetForTest exposes the logout GET handler for tests.
func LogoutGetForTest(w http.ResponseWriter, r *http.Request) { LogoutGet(w, r) }
@@ -463,6 +490,19 @@ func RequireModelAccessForTest(w http.ResponseWriter, r *http.Request, model, pe
// ValidateLoginCSRFForTest exposes pre-session login CSRF validation for tests.
func ValidateLoginCSRFForTest(r *http.Request) bool { return validateLoginCSRF(r) }
+// MintSignedUIDCookieForTest builds a signed uid cookie value for tests.
+func MintSignedUIDCookieForTest(uid int, ttl time.Duration) string {
+ return mintSignedUIDCookieValue(uid, ttl)
+}
+
+// ParseSignedUIDCookieForTest parses a signed uid cookie value for tests.
+func ParseSignedUIDCookieForTest(value string) (int, bool) {
+ return parseSignedUIDCookie(value)
+}
+
+// RateLimitedPathForTest reports whether a path is subject to IP rate limiting.
+func RateLimitedPathForTest(path string) bool { return rateLimitedPath(path) }
+
// TestLoginCSRFCookie is the HttpOnly login CSRF cookie name.
const TestLoginCSRFCookie = loginCSRFCookie
@@ -551,12 +591,22 @@ func NewBusHubForTest() *BusHubForTest {
type SwcBusClientForTest struct{ client *swcBusClient }
func NewSwcBusClientForTest(uid int, buffer int) *SwcBusClientForTest {
- return &SwcBusClientForTest{client: &swcBusClient{uid: uid, send: make(chan []byte, buffer)}}
+ return &SwcBusClientForTest{client: &swcBusClient{
+ uid: uid, send: make(chan []byte, buffer), channels: make(map[string]struct{}),
+ }}
}
func (h *BusHubForTest) Register(c *SwcBusClientForTest) { h.hub.register(c.client) }
-func (h *BusHubForTest) Broadcast(actor int, msg []byte) { h.hub.broadcast(actor, msg) }
+func (h *BusHubForTest) PublishChannel(channel string, msg []byte) {
+ h.hub.publishChannel(channel, msg)
+}
+
+func (c *SwcBusClientForTest) SubscribeChannel(channel string) {
+ c.client.subMu.Lock()
+ c.client.channels[channel] = struct{}{}
+ c.client.subMu.Unlock()
+}
func (c *SwcBusClientForTest) Recv() <-chan []byte { return c.client.send }
diff --git a/core/server/web/web_constants.go b/core/server/web/web_constants.go
index 9486b7b0..5ebaae54 100644
--- a/core/server/web/web_constants.go
+++ b/core/server/web/web_constants.go
@@ -61,13 +61,23 @@ const (
// Login page identifiers and form fields.
const (
- loginTemplateFile = "login.html"
- loginField = "login"
- passwordField = "password"
- nextField = "next"
- invalidLoginMessage = "Invalid login or password."
- resetPasswordMsg = "reset_requested"
- resetUserIDField = "id"
+ loginTemplateFile = "login.html"
+ totpLoginTemplateFile = "totp_login.html"
+ totpLoginRoute = "/web/login/totp"
+ loginCSRFCookie = "sumeru_login_csrf"
+ loginNextCookie = "sumeru_login_next"
+ pendingMFACookie = "sumeru_pending_mfa"
+ trustedDeviceCookie = "sumeru_trusted_device"
+ loginLockoutMaxFailures = 5
+ loginField = "login"
+ passwordField = "password"
+ nextField = "next"
+ invalidLoginMessage = "Invalid login or password."
+ resetPasswordMsg = "reset_requested"
+ oauthDeniedMsg = "oauth_denied"
+ authLocalDisabledMsg = "auth_local_disabled"
+ authLocalConfigKey = "auth.local_enabled"
+ resetUserIDField = "id"
)
// Auth HTTP headers.
diff --git a/core/swc/src/login/password-toggle.ts b/core/swc/src/login/password-toggle.ts
index 859a6901..186099f4 100644
--- a/core/swc/src/login/password-toggle.ts
+++ b/core/swc/src/login/password-toggle.ts
@@ -15,20 +15,10 @@ function setVisible(input: HTMLInputElement, btn: HTMLButtonElement, show: boole
btn.classList.toggle("sum-password-toggle--revealed", show);
}
-function enhanceInput(input: HTMLInputElement): void {
- if (input.closest(".sum-password-field")) {
+function attachToggle(input: HTMLInputElement, wrapper: HTMLElement): void {
+ if (wrapper.querySelector(".sum-password-toggle")) {
return;
}
- const parent = input.parentNode;
- if (!parent) {
- return;
- }
-
- const wrapper = document.createElement("div");
- wrapper.className = "sum-password-field";
- parent.insertBefore(wrapper, input);
- wrapper.appendChild(input);
-
const btn = document.createElement("button");
btn.type = "button";
btn.className = "sum-password-toggle";
@@ -36,12 +26,29 @@ function enhanceInput(input: HTMLInputElement): void {
btn.setAttribute("aria-pressed", "false");
btn.innerHTML = EYE_OPEN_SVG + EYE_CLOSED_SVG;
wrapper.appendChild(btn);
-
btn.addEventListener("click", () => {
setVisible(input, btn, input.type === "password");
});
}
+function enhanceInput(input: HTMLInputElement): void {
+ const existingWrapper = input.closest(".sum-password-field");
+ if (existingWrapper instanceof HTMLElement) {
+ attachToggle(input, existingWrapper);
+ return;
+ }
+ const parent = input.parentNode;
+ if (!parent) {
+ return;
+ }
+
+ const wrapper = document.createElement("div");
+ wrapper.className = "sum-password-field";
+ parent.insertBefore(wrapper, input);
+ wrapper.appendChild(input);
+ attachToggle(input, wrapper);
+}
+
/** Enhance password inputs inside root with show/hide toggles. */
export function initPasswordToggles(root: ParentNode = document): void {
root.querySelectorAll
('input[type="password"]').forEach(enhanceInput);
diff --git a/core/swc/src/main.ts b/core/swc/src/main.ts
index e4a2d367..927a4d7c 100644
--- a/core/swc/src/main.ts
+++ b/core/swc/src/main.ts
@@ -126,7 +126,7 @@ function bootstrap(): void {
initDebugManager(boot, { dialog: env.services.dialog, notification: env.services.notification });
mountDebugEnvironment(boot);
void updateDebugDrawer(boot);
- initShellChrome(boot, env.services.http);
+ initShellChrome(boot, env.services.http, env.services.bus);
initAppLauncher(boot, env.services.action, env.services.command);
const mountEl = document.getElementById("swc-workspace");
diff --git a/core/swc/src/model/record.ts b/core/swc/src/model/record.ts
index 07eb27c5..8d5d55f9 100644
--- a/core/swc/src/model/record.ts
+++ b/core/swc/src/model/record.ts
@@ -155,6 +155,8 @@ export class RecordService {
// ponytail: FIFO eviction at 64 entries; upgrade to LRU if profiling shows churn.
private static readonly maxCache = 64;
+ private readonly recordWatches = new Map void>();
+
constructor(rpc: RpcService, bus: BusService) {
this.store = new RecordStore(rpc);
this.bus = bus;
@@ -186,10 +188,21 @@ export class RecordService {
const rec = this.store.fromPayload(model, id, data);
if (id > 0) {
this.remember(model, id, rec);
+ this.ensureServerWatch(model, id);
}
return rec;
}
+ private ensureServerWatch(model: string, id: number): void {
+ const key = this.cacheKey(model, id);
+ if (this.recordWatches.has(key)) return;
+ const unsub = this.bus.watchRecord(model, id, () => {
+ this.invalidate(model, id);
+ this.bus.emit(RECORD_UPDATED, { model, id, recordId: id });
+ });
+ this.recordWatches.set(key, unsub);
+ }
+
get(model: string, id: number): SwcRecord | undefined {
if (id <= 0) return undefined;
return this.cache.get(this.cacheKey(model, id));
diff --git a/core/swc/src/services/bus.ts b/core/swc/src/services/bus.ts
index 4329e2b0..20e72f6b 100644
--- a/core/swc/src/services/bus.ts
+++ b/core/swc/src/services/bus.ts
@@ -2,17 +2,51 @@ import { SWC_API_BASE } from "../constants/routes.js";
type BusHandler = (payload: unknown) => void;
-function parseBusMessage(raw: unknown): { channel: string; payload: unknown } | null {
+interface BusEventFrame {
+ type?: string;
+ id?: number;
+ channel?: string;
+ payload?: unknown;
+}
+
+function parseBusEventFrame(raw: unknown): BusEventFrame | null {
if (typeof raw !== "object" || raw === null) return null;
- const channel = (raw as { channel?: unknown }).channel;
- if (typeof channel !== "string" || channel === "") return null;
- return { channel, payload: (raw as { payload?: unknown }).payload };
+ return raw as BusEventFrame;
+}
+
+function lastEventStorageKey(channel: string): string {
+ return `sum-bus-last:${channel}`;
+}
+
+function readLastEventId(channel: string): number {
+ try {
+ const v = sessionStorage.getItem(lastEventStorageKey(channel));
+ if (!v) return 0;
+ const n = Number(v);
+ return Number.isFinite(n) && n > 0 ? n : 0;
+ } catch {
+ return 0;
+ }
+}
+
+function writeLastEventId(channel: string, id: number): void {
+ if (id <= 0) return;
+ try {
+ sessionStorage.setItem(lastEventStorageKey(channel), String(id));
+ } catch {
+ /* ignore */
+ }
}
-/** Client event bus with optional WebSocket live updates from /web/swc/bus. */
+/** Client event bus with WebSocket live updates from /web/swc/bus. */
export class BusService {
private readonly handlers = new Map>();
private ws: WebSocket | null = null;
+ private wsURL = `${SWC_API_BASE}/bus`;
+ private reconnectAttempt = 0;
+ private reconnectTimer: ReturnType | null = null;
+ private remoteChannels = new Set();
+ private intentionalClose = false;
subscribe(channel: string, handler: BusHandler): () => void {
if (!this.handlers.has(channel)) {
@@ -28,31 +62,105 @@ export class BusService {
}
}
- /** Connect to /web/swc/bus when bootstrap.busEnabled is true. */
+ /** Subscribe on server when connected; always registers local handler via subscribe(). */
+ watchChannel(channel: string, handler: BusHandler): () => void {
+ const unsubLocal = this.subscribe(channel, handler);
+ this.remoteChannels.add(channel);
+ this.sendSubscribe([channel]);
+ return () => {
+ unsubLocal();
+ this.remoteChannels.delete(channel);
+ this.sendUnsubscribe([channel]);
+ };
+ }
+
+ watchRecord(model: string, id: number, handler: BusHandler): () => void {
+ if (!model || id <= 0) return () => undefined;
+ const channel = `record/${model}/${id}`;
+ return this.watchChannel(channel, handler);
+ }
+
connect(url = `${SWC_API_BASE}/bus`): void {
+ this.wsURL = url;
+ this.intentionalClose = false;
+ this.openSocket();
+ }
+
+ disconnect(): void {
+ this.intentionalClose = true;
+ if (this.reconnectTimer) {
+ clearTimeout(this.reconnectTimer);
+ this.reconnectTimer = null;
+ }
+ this.ws?.close();
+ this.ws = null;
+ }
+
+ private openSocket(): void {
if (this.ws) return;
try {
const proto = window.location.protocol === "https:" ? "wss:" : "ws:";
- this.ws = new WebSocket(`${proto}//${window.location.host}${url}`);
+ this.ws = new WebSocket(`${proto}//${window.location.host}${this.wsURL}`);
+ this.ws.addEventListener("open", () => {
+ this.reconnectAttempt = 0;
+ if (this.remoteChannels.size > 0) {
+ this.sendSubscribe([...this.remoteChannels]);
+ }
+ });
this.ws.addEventListener("message", (ev) => {
try {
const parsed: unknown = JSON.parse(String(ev.data));
- const msg = parseBusMessage(parsed);
- if (msg) this.emit(msg.channel, msg.payload);
+ const frame = parseBusEventFrame(parsed);
+ if (!frame) return;
+ if (frame.type === "event" && frame.channel) {
+ if (typeof frame.id === "number" && frame.id > 0) {
+ writeLastEventId(frame.channel, frame.id);
+ }
+ this.emit(frame.channel, frame.payload);
+ }
} catch (err) {
console.warn("swc bus: malformed message", err);
}
});
this.ws.addEventListener("close", () => {
this.ws = null;
+ if (!this.intentionalClose) {
+ this.scheduleReconnect();
+ }
});
} catch (err) {
console.warn("swc bus: WebSocket unavailable; local-only bus", err);
+ this.scheduleReconnect();
}
}
- disconnect(): void {
- this.ws?.close();
- this.ws = null;
+ private scheduleReconnect(): void {
+ if (this.reconnectTimer || this.intentionalClose) return;
+ const delay = Math.min(60_000, 1000 * 2 ** this.reconnectAttempt) * (0.8 + Math.random() * 0.4);
+ this.reconnectAttempt += 1;
+ this.reconnectTimer = setTimeout(() => {
+ this.reconnectTimer = null;
+ this.openSocket();
+ }, delay);
+ }
+
+ private sendFrame(obj: Record): void {
+ if (!this.ws || this.ws.readyState !== WebSocket.OPEN) return;
+ this.ws.send(JSON.stringify(obj));
+ }
+
+ private sendSubscribe(channels: string[]): void {
+ if (channels.length === 0) return;
+ const frame: Record = { type: "subscribe", channels };
+ const maxLast = Math.max(0, ...channels.map((c) => readLastEventId(c)));
+ if (maxLast > 0) {
+ frame.last_event_id = maxLast;
+ }
+ this.sendFrame(frame);
+ }
+
+ private sendUnsubscribe(channels: string[]): void {
+ if (channels.length === 0) return;
+ this.sendFrame({ type: "unsubscribe", channels });
}
}
diff --git a/core/swc/src/shell/notification-bell.ts b/core/swc/src/shell/notification-bell.ts
new file mode 100644
index 00000000..42548d6f
--- /dev/null
+++ b/core/swc/src/shell/notification-bell.ts
@@ -0,0 +1,80 @@
+import type { BusService } from "../services/bus.js";
+import type { HttpService } from "../services/http.js";
+import { SWC_API_BASE } from "../constants/routes.js";
+
+interface InboxPayload {
+ items: Array<{
+ id: number;
+ body: string;
+ author: string;
+ recordModel: string;
+ recordId: number;
+ isRead: boolean;
+ }>;
+ unread: number;
+}
+
+export function initNotificationBell(http: HttpService, bus: BusService, userId: number): void {
+ const host = document.querySelector(".sum-topbar-right");
+ if (!host || userId <= 0) return;
+
+ const btn = document.createElement("button");
+ btn.type = "button";
+ btn.className = "sum-icon-btn sum-notification-bell";
+ btn.setAttribute("aria-label", "Notifications");
+ btn.innerHTML = `●`;
+
+ const panel = document.createElement("div");
+ panel.className = "sum-notification-panel";
+ panel.hidden = true;
+
+ host.insertBefore(btn, host.firstChild);
+ host.insertBefore(panel, btn.nextSibling);
+
+ const badge = btn.querySelector(".sum-notification-badge") as HTMLElement;
+ const base = SWC_API_BASE;
+
+ async function refresh(): Promise {
+ try {
+ const data = await http.getJSON(`${base}/notifications?unread=0`);
+ badge.textContent = data.unread > 0 ? String(data.unread) : "";
+ badge.hidden = data.unread <= 0;
+ panel.innerHTML = "";
+ for (const item of data.items.slice(0, 20)) {
+ const row = document.createElement("button");
+ row.type = "button";
+ row.className = "sum-notification-row" + (item.isRead ? "" : " sum-notification-row--unread");
+ row.textContent = `${item.author}: ${item.body}`.slice(0, 120);
+ row.addEventListener("click", () => {
+ void http.postForm(`${base}/notifications/read`, { id: String(item.id) });
+ if (item.recordModel && item.recordId > 0) {
+ window.location.assign(`/web#action=&model=${encodeURIComponent(item.recordModel)}&id=${item.recordId}`);
+ }
+ void refresh();
+ });
+ panel.appendChild(row);
+ }
+ const markAll = document.createElement("button");
+ markAll.type = "button";
+ markAll.className = "sum-notification-mark-all";
+ markAll.textContent = "Mark all read";
+ markAll.addEventListener("click", () => {
+ void http.postForm(`${base}/notifications/read`, { all: "1" }).then(refresh);
+ });
+ panel.appendChild(markAll);
+ } catch {
+ /* ignore */
+ }
+ }
+
+ btn.addEventListener("click", () => {
+ panel.hidden = !panel.hidden;
+ if (!panel.hidden) void refresh();
+ });
+
+ bus.watchChannel(`user/${userId}/notifications`, () => {
+ void refresh();
+ });
+
+ void refresh();
+}
diff --git a/core/swc/src/shell/shell-chrome.ts b/core/swc/src/shell/shell-chrome.ts
index 5f82e141..c8ed5c9e 100644
--- a/core/swc/src/shell/shell-chrome.ts
+++ b/core/swc/src/shell/shell-chrome.ts
@@ -8,8 +8,10 @@ import { initSidebar } from "./sidebar.js";
import { initCompanySwitcher } from "./company-switcher.js";
import { initViewTabNavigation } from "./view-tab-sync.js";
import { initBreadcrumbNavigation } from "./breadcrumb-sync.js";
+import { initNotificationBell } from "./notification-bell.js";
+import type { BusService } from "../services/bus.js";
-export function initShellChrome(boot: SwcBootstrap, http: HttpService): void {
+export function initShellChrome(boot: SwcBootstrap, http: HttpService, bus?: BusService): void {
const shell = document.getElementById("sum-shell");
if (!shell) return;
@@ -25,5 +27,9 @@ export function initShellChrome(boot: SwcBootstrap, http: HttpService): void {
initHomeDashboard(http);
initCompanySwitcher(boot, http);
+ if (bus && boot.user?.id) {
+ initNotificationBell(http, bus, boot.user.id);
+ }
+
new NotificationService().bootstrap(boot.toasts);
}
diff --git a/core/swc/tests/login/password-toggle.test.ts b/core/swc/tests/login/password-toggle.test.ts
index a5db58e9..ded12802 100644
--- a/core/swc/tests/login/password-toggle.test.ts
+++ b/core/swc/tests/login/password-toggle.test.ts
@@ -46,4 +46,14 @@ describe("initPasswordToggles", () => {
initPasswordToggles(document);
expect(document.querySelectorAll(".sum-password-field").length).toBe(2);
});
+
+ it("adds toggle to pre-wrapped login-style password fields", () => {
+ document.body.innerHTML =
+ '';
+ initPasswordToggles(document);
+ const wrapper = document.querySelector(".sum-password-field");
+ expect(wrapper?.querySelectorAll(".sum-password-toggle").length).toBe(1);
+ initPasswordToggles(document);
+ expect(wrapper?.querySelectorAll(".sum-password-toggle").length).toBe(1);
+ });
});
diff --git a/go.mod b/go.mod
index 53290842..e89e88c4 100644
--- a/go.mod
+++ b/go.mod
@@ -11,6 +11,8 @@ require (
gopkg.in/natefinch/lumberjack.v2 v2.2.1
)
+require github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
+
replace github.com/gpdf-dev/gpdf => github.com/ProjectMeru/gpdf v1.0.13
replace github.com/gorilla/websocket => github.com/ProjectMeru/websocket v1.5.3
diff --git a/go.sum b/go.sum
index 1c0d42b4..15e5cbeb 100644
--- a/go.sum
+++ b/go.sum
@@ -7,6 +7,8 @@ github.com/ProjectMeru/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+E
github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE=
github.com/lib/pq v1.12.3 h1:tTWxr2YLKwIvK90ZXEw8GP7UFHtcbTtty8zsI+YjrfQ=
github.com/lib/pq v1.12.3/go.mod h1:/p+8NSbOcwzAEI7wiMXFlgydTwcgTr3OSKMsD2BitpA=
+github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e h1:MRM5ITcdelLK2j1vwZ3Je0FKVCfqOLp5zO6trqMLYs0=
+github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e/go.mod h1:XV66xRDqSt+GTGFMVlhk3ULuV0y9ZmzeVGR4mloJI3M=
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc=
diff --git a/test/core/coverage/exports_test.go b/test/core/coverage/exports_test.go
index 5bcaeef4..7b72ee20 100644
--- a/test/core/coverage/exports_test.go
+++ b/test/core/coverage/exports_test.go
@@ -353,7 +353,8 @@ func TestForTestExports_web(t *testing.T) {
hub := web.NewBusHubForTest()
client := web.NewSwcBusClientForTest(1, 2)
hub.Register(client)
- hub.Broadcast(1, []byte("ping"))
+ client.SubscribeChannel("user/1/notifications")
+ hub.PublishChannel("user/1/notifications", []byte("ping"))
select {
case msg := <-client.Recv():
if string(msg) != "ping" {
diff --git a/test/core/orm/sqlmock_test.go b/test/core/orm/sqlmock_test.go
index bfeb44e2..fa231b8d 100644
--- a/test/core/orm/sqlmock_test.go
+++ b/test/core/orm/sqlmock_test.go
@@ -31,6 +31,20 @@ func bypassCtx() context.Context {
return orm.ContextWithBypass(context.Background(), true)
}
+func TestGetConfig_readsValue(t *testing.T) {
+ mock := setupMockORM(t)
+ mock.ExpectQuery(`SELECT value FROM .+ WHERE key = \$1`).
+ WithArgs("auth.local_enabled").
+ WillReturnRows(sqlmock.NewRows([]string{"value"}).AddRow("false"))
+ got := orm.GetConfig(bypassCtx(), "auth.local_enabled", "true")
+ if got != "false" {
+ t.Fatalf("got %q", got)
+ }
+ if err := mock.ExpectationsWereMet(); err != nil {
+ t.Fatal(err)
+ }
+}
+
func TestSearchWithMockDB(t *testing.T) {
mock := setupMockORM(t)
rows := sqlmock.NewRows([]string{"id", "name", "active"}).
diff --git a/test/core/orm/user_totp_test.go b/test/core/orm/user_totp_test.go
new file mode 100644
index 00000000..75b99596
--- /dev/null
+++ b/test/core/orm/user_totp_test.go
@@ -0,0 +1,61 @@
+package orm_test
+
+import (
+ "context"
+ "strings"
+ "testing"
+
+ "sumeru/core/orm"
+)
+
+func TestBeginOwnTOTPEnrollment_requiresAuth(t *testing.T) {
+ _, err := orm.BeginOwnTOTPEnrollment(context.Background(), 0, 0)
+ if err == nil {
+ t.Fatal("expected error")
+ }
+}
+
+func TestBeginOwnTOTPEnrollment_requiresOwnAccount(t *testing.T) {
+ ctx := orm.ContextWithUID(context.Background(), 2)
+ _, err := orm.BeginOwnTOTPEnrollment(ctx, 2, 3)
+ if err == nil || !strings.Contains(err.Error(), "own account") {
+ t.Fatalf("err=%v", err)
+ }
+}
+
+func TestDisableUserTOTP_requiresAdmin(t *testing.T) {
+ ctx := orm.ContextWithUID(context.Background(), 2)
+ err := orm.DisableUserTOTP(ctx, 2, 3)
+ if err == nil || !strings.Contains(err.Error(), "system administrator") {
+ t.Fatalf("err=%v", err)
+ }
+}
+
+func TestConfirmOwnTOTPEnrollment_requiresAuth(t *testing.T) {
+ err := orm.ConfirmOwnTOTPEnrollment(context.Background(), 0, 0, "123456")
+ if err == nil {
+ t.Fatal("expected error")
+ }
+}
+
+func TestDisableOwnTOTP_requiresAuth(t *testing.T) {
+ err := orm.DisableOwnTOTP(context.Background(), 0, 0, "123456")
+ if err == nil {
+ t.Fatal("expected error")
+ }
+}
+
+func TestUserTOTPCredentialsForLogin_withoutDB(t *testing.T) {
+ secret, enabled := orm.UserTOTPCredentialsForLogin(context.Background(), 1)
+ if secret != "" || enabled {
+ t.Fatalf("secret=%q enabled=%v", secret, enabled)
+ }
+ _, ok, err := orm.OwnTOTPEnrollmentPending(context.Background(), 1, 1)
+ if err == nil || ok {
+ t.Fatalf("ok=%v err=%v", ok, err)
+ }
+ enabled, err = orm.UserTOTPEnabled(context.Background(), 1)
+ if err == nil || enabled {
+ t.Fatalf("enabled=%v err=%v", enabled, err)
+ }
+}
diff --git a/test/core/server/auth/oidc_verify_test.go b/test/core/server/auth/oidc_verify_test.go
new file mode 100644
index 00000000..4ddea3a2
--- /dev/null
+++ b/test/core/server/auth/oidc_verify_test.go
@@ -0,0 +1,199 @@
+package auth_test
+
+import (
+ "context"
+ "crypto"
+ "crypto/rand"
+ "crypto/rsa"
+ "crypto/sha256"
+ "encoding/base64"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "sumeru/core/server/auth"
+)
+
+func TestVerifyIDToken_expired(t *testing.T) {
+ key, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatal(err)
+ }
+ n := base64.RawURLEncoding.EncodeToString(key.N.Bytes())
+ e := base64.RawURLEncoding.EncodeToString([]byte{0x01, 0x00, 0x01})
+ jwks := `{"keys":[{"kty":"RSA","kid":"t1","use":"sig","alg":"RS256","n":"` + n + `","e":"` + e + `"}]}`
+ srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ _, _ = w.Write([]byte(jwks))
+ }))
+ defer srv.Close()
+ header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","kid":"t1","typ":"JWT"}`))
+ payload, _ := json.Marshal(map[string]interface{}{
+ "iss": "https://idp.example",
+ "sub": "user-1",
+ "aud": "client-abc",
+ "exp": time.Now().Add(-time.Hour).Unix(),
+ })
+ payloadB64 := base64.RawURLEncoding.EncodeToString(payload)
+ signingInput := header + "." + payloadB64
+ digest := sha256.Sum256([]byte(signingInput))
+ sig, err := rsa.SignPKCS1v15(rand.Reader, key, crypto.SHA256, digest[:])
+ if err != nil {
+ t.Fatal(err)
+ }
+ token := signingInput + "." + base64.RawURLEncoding.EncodeToString(sig)
+ _, err = auth.VerifyIDToken(context.Background(), token, "https://idp.example", "client-abc", srv.URL)
+ if err == nil {
+ t.Fatal("expected expired error")
+ }
+}
+
+func TestVerifyIDToken_wrongIssuer(t *testing.T) {
+ key, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatal(err)
+ }
+ n := base64.RawURLEncoding.EncodeToString(key.N.Bytes())
+ e := base64.RawURLEncoding.EncodeToString([]byte{0x01, 0x00, 0x01})
+ jwks := `{"keys":[{"kty":"RSA","kid":"t1","use":"sig","alg":"RS256","n":"` + n + `","e":"` + e + `"}]}`
+ srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ _, _ = w.Write([]byte(jwks))
+ }))
+ defer srv.Close()
+ header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","kid":"t1","typ":"JWT"}`))
+ payload, _ := json.Marshal(map[string]interface{}{
+ "iss": "https://wrong.example",
+ "sub": "user-1",
+ "aud": "client-abc",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+ payloadB64 := base64.RawURLEncoding.EncodeToString(payload)
+ signingInput := header + "." + payloadB64
+ digest := sha256.Sum256([]byte(signingInput))
+ sig, err := rsa.SignPKCS1v15(rand.Reader, key, crypto.SHA256, digest[:])
+ if err != nil {
+ t.Fatal(err)
+ }
+ token := signingInput + "." + base64.RawURLEncoding.EncodeToString(sig)
+ _, err = auth.VerifyIDToken(context.Background(), token, "https://idp.example", "client-abc", srv.URL)
+ if err == nil {
+ t.Fatal("expected issuer mismatch")
+ }
+}
+
+func TestVerifyIDToken_wrongAudience(t *testing.T) {
+ key, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatal(err)
+ }
+ n := base64.RawURLEncoding.EncodeToString(key.N.Bytes())
+ e := base64.RawURLEncoding.EncodeToString([]byte{0x01, 0x00, 0x01})
+ jwks := `{"keys":[{"kty":"RSA","kid":"t1","use":"sig","alg":"RS256","n":"` + n + `","e":"` + e + `"}]}`
+ srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ _, _ = w.Write([]byte(jwks))
+ }))
+ defer srv.Close()
+ header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","kid":"t1","typ":"JWT"}`))
+ payload, _ := json.Marshal(map[string]interface{}{
+ "iss": "https://idp.example",
+ "sub": "user-1",
+ "aud": "other-client",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+ payloadB64 := base64.RawURLEncoding.EncodeToString(payload)
+ signingInput := header + "." + payloadB64
+ digest := sha256.Sum256([]byte(signingInput))
+ sig, err := rsa.SignPKCS1v15(rand.Reader, key, crypto.SHA256, digest[:])
+ if err != nil {
+ t.Fatal(err)
+ }
+ token := signingInput + "." + base64.RawURLEncoding.EncodeToString(sig)
+ _, err = auth.VerifyIDToken(context.Background(), token, "https://idp.example", "client-abc", srv.URL)
+ if err == nil {
+ t.Fatal("expected audience mismatch")
+ }
+}
+
+func TestVerifyIDToken_rs256(t *testing.T) {
+ key, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatal(err)
+ }
+ n := base64.RawURLEncoding.EncodeToString(key.N.Bytes())
+ e := base64.RawURLEncoding.EncodeToString([]byte{0x01, 0x00, 0x01})
+ jwks := `{"keys":[{"kty":"RSA","kid":"t1","use":"sig","alg":"RS256","n":"` + n + `","e":"` + e + `"}]}`
+ srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ _, _ = w.Write([]byte(jwks))
+ }))
+ defer srv.Close()
+
+ header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","kid":"t1","typ":"JWT"}`))
+ payload, _ := json.Marshal(map[string]interface{}{
+ "iss": "https://idp.example",
+ "sub": "user-1",
+ "aud": []interface{}{"client-abc", "other"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ "email": "a@example.com",
+ })
+ payloadB64 := base64.RawURLEncoding.EncodeToString(payload)
+ signingInput := header + "." + payloadB64
+ digest := sha256.Sum256([]byte(signingInput))
+ sig, err := rsa.SignPKCS1v15(rand.Reader, key, crypto.SHA256, digest[:])
+ if err != nil {
+ t.Fatal(err)
+ }
+ token := signingInput + "." + base64.RawURLEncoding.EncodeToString(sig)
+
+ claims, err := auth.VerifyIDToken(context.Background(), token, "https://idp.example", "client-abc", srv.URL)
+ if err != nil {
+ t.Fatalf("verify: %v", err)
+ }
+ if claims["sub"] != "user-1" {
+ t.Fatalf("sub=%v", claims["sub"])
+ }
+}
+
+func TestParseIDTokenClaimsUnsafe(t *testing.T) {
+ email, sub, claims := auth.ParseIDTokenClaimsUnsafe("")
+ if email != "" || sub != "" || claims != nil {
+ t.Fatal("empty token should yield nothing")
+ }
+ payload := base64.RawURLEncoding.EncodeToString([]byte(`{"sub":"u1","email":"a@example.com"}`))
+ token := "e30." + payload + ".sig"
+ email, sub, claims = auth.ParseIDTokenClaimsUnsafe(token)
+ if email != "a@example.com" || sub != "u1" || claims == nil {
+ t.Fatalf("email=%q sub=%q claims=%v", email, sub, claims)
+ }
+}
+
+func TestVerifyIDToken_rejectsMalformed(t *testing.T) {
+ _, err := auth.VerifyIDToken(context.Background(), "not.a.jwt", "", "", "http://127.0.0.1")
+ if err == nil {
+ t.Fatal("expected error")
+ }
+}
+
+func TestVerifyIDToken_rejectsUnsupportedAlg(t *testing.T) {
+ header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"HS256","typ":"JWT"}`))
+ payload := base64.RawURLEncoding.EncodeToString([]byte(`{"sub":"x"}`))
+ token := header + "." + payload + ".sig"
+ _, err := auth.VerifyIDToken(context.Background(), token, "", "", "http://127.0.0.1/jwks")
+ if err == nil {
+ t.Fatal("expected unsupported alg error")
+ }
+}
+
+func TestTOTPOtpauthURI(t *testing.T) {
+ uri := auth.TOTPOtpauthURI("Sumeru", "admin@example.com", "JBSWY3DPEHPK3PXP")
+ if uri == "" || len(uri) < 20 {
+ t.Fatalf("uri=%q", uri)
+ }
+}
+
+func TestTOTPQRPng(t *testing.T) {
+ png, err := auth.TOTPQRPng(auth.TOTPOtpauthURI("Sumeru", "u", "JBSWY3DPEHPK3PXP"), 120)
+ if err != nil || len(png) < 32 {
+ t.Fatalf("png len=%d err=%v", len(png), err)
+ }
+}
diff --git a/test/core/server/auth/totp_test.go b/test/core/server/auth/totp_test.go
new file mode 100644
index 00000000..217496f0
--- /dev/null
+++ b/test/core/server/auth/totp_test.go
@@ -0,0 +1,26 @@
+package auth_test
+
+import (
+ "testing"
+
+ "sumeru/core/server/auth"
+)
+
+func TestValidateTOTP_rejectsEmptySecret(t *testing.T) {
+ if auth.ValidateTOTP("", "123456") {
+ t.Fatal("empty secret should fail")
+ }
+}
+
+func TestValidateTOTP_window(t *testing.T) {
+ secret, err := auth.GenerateTOTPSecret()
+ if err != nil {
+ t.Fatal(err)
+ }
+ if auth.ValidateTOTP(secret, "abc") {
+ t.Fatal("non-numeric code should fail")
+ }
+ if auth.ValidateTOTP(secret, "000000") && auth.ValidateTOTP(secret, "111111") {
+ t.Fatal("random secret should not match two arbitrary codes")
+ }
+}
diff --git a/test/core/server/web/login_csrf_test.go b/test/core/server/web/login_csrf_test.go
index e19d0b40..f551de1e 100644
--- a/test/core/server/web/login_csrf_test.go
+++ b/test/core/server/web/login_csrf_test.go
@@ -161,6 +161,41 @@ func hiddenInputValue(html, name string) string {
return m[1]
}
+func TestLoginGet_showsInfoFlashFromQuery(t *testing.T) {
+ root := sumeruModuleRoot(t)
+ prevTemplates := config.AppConfig.TemplatesPath
+ config.AppConfig.TemplatesPath = filepath.Join(root, "core", "engine", "templates")
+ t.Cleanup(func() {
+ config.AppConfig.TemplatesPath = prevTemplates
+ })
+
+ req := httptest.NewRequest(http.MethodGet, web.TestLoginRoute+"?msg="+web.TestResetPasswordMsg, nil)
+ rec := httptest.NewRecorder()
+ web.LoginGetForTest(rec, req)
+ if !strings.Contains(rec.Body.String(), "sum-login-info") {
+ t.Fatal("expected info flash region")
+ }
+}
+
+func TestLoginGet_rendersEnterpriseShell(t *testing.T) {
+ root := sumeruModuleRoot(t)
+ prevTemplates := config.AppConfig.TemplatesPath
+ config.AppConfig.TemplatesPath = filepath.Join(root, "core", "engine", "templates")
+ t.Cleanup(func() {
+ config.AppConfig.TemplatesPath = prevTemplates
+ })
+
+ req := httptest.NewRequest(http.MethodGet, web.TestLoginRoute+"?msg="+web.TestOAuthDeniedMsg, nil)
+ rec := httptest.NewRecorder()
+ web.LoginGetForTest(rec, req)
+ body := rec.Body.String()
+ for _, needle := range []string{"sum-login-shell", "sum-login-brand", "Sign in", "Authorized users only"} {
+ if !strings.Contains(body, needle) {
+ t.Fatalf("missing %q in login HTML", needle)
+ }
+ }
+}
+
func sumeruModuleRoot(t *testing.T) string {
t.Helper()
dir, err := os.Getwd()
diff --git a/test/core/server/web/oauth_callback_test.go b/test/core/server/web/oauth_callback_test.go
new file mode 100644
index 00000000..fa9467fb
--- /dev/null
+++ b/test/core/server/web/oauth_callback_test.go
@@ -0,0 +1,46 @@
+package web_test
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "sumeru/core/server/web"
+)
+
+func TestOAuthCallback_invalidStateRedirectsToLogin(t *testing.T) {
+ req := httptest.NewRequest(http.MethodGet, "/web/auth/oauth/callback?state=bad&code=abc", nil)
+ rec := httptest.NewRecorder()
+ web.OAuthCallbackForTest(rec, req)
+ if rec.Code != http.StatusFound {
+ t.Fatalf("status=%d", rec.Code)
+ }
+ loc := rec.Header().Get("Location")
+ if !strings.Contains(loc, web.TestOAuthDeniedMsg) {
+ t.Fatalf("location=%q", loc)
+ }
+}
+
+func TestOAuthStart_missingProvider(t *testing.T) {
+ req := httptest.NewRequest(http.MethodGet, "/web/auth/oauth/start", nil)
+ rec := httptest.NewRecorder()
+ web.OAuthStartForTest(rec, req)
+ if rec.Code != http.StatusBadRequest {
+ t.Fatalf("status=%d", rec.Code)
+ }
+}
+
+func TestAuthLocalEnabled_defaultsTrue(t *testing.T) {
+ if !web.AuthLocalEnabledForTest(context.Background()) {
+ t.Fatal("expected local login enabled by default")
+ }
+}
+
+func TestListEnabledAuthProviders_emptyWithoutDB(t *testing.T) {
+ if got := web.ListEnabledAuthProvidersForTest(context.Background()); len(got) != 0 {
+ t.Fatalf("expected no providers, got %v", got)
+ }
+ _ = web.LoginPageCompanyNameForTest(context.Background())
+}
diff --git a/test/core/server/web/pending_mfa_cookie_test.go b/test/core/server/web/pending_mfa_cookie_test.go
new file mode 100644
index 00000000..979407f0
--- /dev/null
+++ b/test/core/server/web/pending_mfa_cookie_test.go
@@ -0,0 +1,24 @@
+package web_test
+
+import (
+ "testing"
+ "time"
+
+ "sumeru/core/server/web"
+)
+
+func TestSignedUIDCookie_roundTrip(t *testing.T) {
+ val := web.MintSignedUIDCookieForTest(42, time.Minute)
+ uid, ok := web.ParseSignedUIDCookieForTest(val)
+ if !ok || uid != 42 {
+ t.Fatalf("parse signed uid: ok=%v uid=%d", ok, uid)
+ }
+}
+
+func TestSignedUIDCookie_rejectsTamper(t *testing.T) {
+ val := web.MintSignedUIDCookieForTest(42, time.Minute)
+ tampered := val[:len(val)-2] + "xx"
+ if uid, ok := web.ParseSignedUIDCookieForTest(tampered); ok || uid != 0 {
+ t.Fatalf("expected tampered cookie rejected, got uid=%d ok=%v", uid, ok)
+ }
+}
diff --git a/test/core/server/web/query_flash_sanitize_test.go b/test/core/server/web/query_flash_sanitize_test.go
new file mode 100644
index 00000000..75401cc4
--- /dev/null
+++ b/test/core/server/web/query_flash_sanitize_test.go
@@ -0,0 +1,28 @@
+package web_test
+
+import (
+ "strings"
+ "testing"
+ "time"
+ "unicode/utf8"
+
+ "sumeru/core/server/web"
+)
+
+func TestFlashFromQueryMessage_truncatesLongErrorBody(t *testing.T) {
+ long := strings.Repeat("x", 600)
+ flash, ok := web.FlashFromQueryMessage("error:" + long)
+ if !ok {
+ t.Fatal("expected flash")
+ }
+ if utf8.RuneCountInString(flash.Body) > 500 {
+ t.Fatalf("body runes=%d want <=500", utf8.RuneCountInString(flash.Body))
+ }
+}
+
+func TestSignedUIDCookie_rejectsExpired(t *testing.T) {
+ val := web.MintSignedUIDCookieForTest(7, -time.Second)
+ if uid, ok := web.ParseSignedUIDCookieForTest(val); ok || uid != 0 {
+ t.Fatalf("expected expired cookie rejected, uid=%d ok=%v", uid, ok)
+ }
+}
diff --git a/test/core/server/web/query_flash_test.go b/test/core/server/web/query_flash_test.go
index 34fb3dab..9f52f4fb 100644
--- a/test/core/server/web/query_flash_test.go
+++ b/test/core/server/web/query_flash_test.go
@@ -43,3 +43,20 @@ func TestFlashFromQuerySaveOKUpdated(t *testing.T) {
t.Fatalf("flash = %+v ok=%v", flash, ok)
}
}
+
+func TestFlashFromQueryAuthMessages(t *testing.T) {
+ cases := []struct {
+ msg, kind, title string
+ }{
+ {web.TestOAuthDeniedMsg, "error", "Sign-in failed"},
+ {web.TestAuthLocalDisabledMsg, "error", "Password sign-in disabled"},
+ {"totp_enabled", "success", "Two-factor enabled"},
+ {"totp_invalid", "error", "Invalid code"},
+ }
+ for _, c := range cases {
+ flash, ok := web.FlashFromQueryMessage(c.msg)
+ if !ok || flash.Kind != c.kind || flash.Title != c.title {
+ t.Fatalf("msg=%q flash=%+v ok=%v", c.msg, flash, ok)
+ }
+ }
+}
diff --git a/test/core/server/web/rate_limit_test.go b/test/core/server/web/rate_limit_test.go
new file mode 100644
index 00000000..0c1120c1
--- /dev/null
+++ b/test/core/server/web/rate_limit_test.go
@@ -0,0 +1,24 @@
+package web_test
+
+import (
+ "testing"
+
+ "sumeru/core/server/web"
+)
+
+func TestRateLimitedPath_authRoutes(t *testing.T) {
+ for _, path := range []string{
+ web.TestLoginRoute,
+ "/web/login/totp",
+ "/web/auth/oauth/start",
+ "/web/auth/oauth/callback",
+ web.TestAPIRPCRoute,
+ } {
+ if !web.RateLimitedPathForTest(path) {
+ t.Fatalf("expected rate limit for %q", path)
+ }
+ }
+ if web.RateLimitedPathForTest("/web/home") {
+ t.Fatal("home should not be rate limited")
+ }
+}
diff --git a/test/core/server/web/swc_bus_channel_test.go b/test/core/server/web/swc_bus_channel_test.go
new file mode 100644
index 00000000..ae3bfe02
--- /dev/null
+++ b/test/core/server/web/swc_bus_channel_test.go
@@ -0,0 +1,25 @@
+package web_test
+
+import (
+ "context"
+ "testing"
+
+ "sumeru/core/server/web"
+)
+
+func TestAuthorizeSwcBusChannel_userScope(t *testing.T) {
+ ctx := context.Background()
+ if web.AuthorizeSwcBusChannel(ctx, 5, "user/5/notifications") != true {
+ t.Fatal("expected own user channel allowed")
+ }
+ if web.AuthorizeSwcBusChannel(ctx, 5, "user/6/notifications") {
+ t.Fatal("expected other user channel denied")
+ }
+}
+
+func TestRecordBusChannel_format(t *testing.T) {
+ ch := web.RecordBusChannel("crm.lead", 9)
+ if ch != "record/crm.lead/9" {
+ t.Fatalf("channel: %q", ch)
+ }
+}
diff --git a/test/core/server/web/swc_bus_hub_test.go b/test/core/server/web/swc_bus_hub_test.go
index cce33bcd..8232c0ad 100644
--- a/test/core/server/web/swc_bus_hub_test.go
+++ b/test/core/server/web/swc_bus_hub_test.go
@@ -7,15 +7,16 @@ import (
"sumeru/core/server/web"
)
-func TestSwcBusHubBroadcastFiltersByActor(t *testing.T) {
+func TestSwcBusHubPublishChannelTargetsSubscribers(t *testing.T) {
h := web.NewBusHubForTest()
userA := web.NewSwcBusClientForTest(1, 1)
userB := web.NewSwcBusClientForTest(2, 1)
h.Register(userA)
h.Register(userB)
+ userA.SubscribeChannel("record/crm.lead/1")
- msg := []byte(`{"channel":"record.updated","payload":{"model":"crm.lead","id":1}}`)
- h.Broadcast(1, msg)
+ msg := []byte(`{"type":"event","id":1,"channel":"record/crm.lead/1","payload":{"model":"crm.lead","id":1}}`)
+ h.PublishChannel("record/crm.lead/1", msg)
select {
case got := <-userA.Recv():
@@ -23,33 +24,28 @@ func TestSwcBusHubBroadcastFiltersByActor(t *testing.T) {
t.Fatalf("user A: got %q", got)
}
default:
- t.Fatal("expected message for user A")
+ t.Fatal("expected message for subscribed user A")
}
select {
case <-userB.Recv():
- t.Fatal("user B should not receive actor-scoped message")
+ t.Fatal("user B should not receive without subscription")
default:
}
}
-func TestSwcBusHubQueueMessageShape(t *testing.T) {
+func TestSwcBusHubEventFrameShape(t *testing.T) {
h := web.NewBusHubForTest()
client := web.NewSwcBusClientForTest(5, 1)
h.Register(client)
+ client.SubscribeChannel("model/core.partner")
- raw, _ := json.Marshal(map[string]interface{}{
- "name": "record.updated",
+ out, _ := json.Marshal(map[string]interface{}{
+ "type": "event",
+ "id": 2,
+ "channel": "model/core.partner",
"payload": map[string]interface{}{"model": "core.partner", "id": 3},
- "actor": 5,
})
- var envelope map[string]interface{}
- if err := json.Unmarshal(raw, &envelope); err != nil {
- t.Fatal(err)
- }
- name := envelope["name"].(string)
- inner := envelope["payload"].(map[string]interface{})
- out, _ := json.Marshal(map[string]interface{}{"channel": name, "payload": inner})
- h.Broadcast(5, out)
+ h.PublishChannel("model/core.partner", out)
select {
case got := <-client.Recv():
@@ -57,7 +53,7 @@ func TestSwcBusHubQueueMessageShape(t *testing.T) {
if err := json.Unmarshal(got, &parsed); err != nil {
t.Fatal(err)
}
- if parsed["channel"] != "record.updated" {
+ if parsed["channel"] != "model/core.partner" {
t.Fatalf("channel: %v", parsed["channel"])
}
default: