From 7771fc0ca490996f8c46eae6407d110d1af51869 Mon Sep 17 00:00:00 2001 From: AIRONAX Developer Date: Fri, 2 Oct 2026 20:29:50 +0530 Subject: [PATCH 01/14] fix(mail): wrap mail subtype data in sumeru root element Addon data must use the sumeru document root so generate and module load succeed. --- addons/mail/data/mail_subtype_data.xml | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) create mode 100644 addons/mail/data/mail_subtype_data.xml diff --git a/addons/mail/data/mail_subtype_data.xml b/addons/mail/data/mail_subtype_data.xml new file mode 100644 index 0000000..dd805b0 --- /dev/null +++ b/addons/mail/data/mail_subtype_data.xml @@ -0,0 +1,16 @@ + + + + + comment + User discussion + True + + + notification + System notification + True + True + + + From 47103a5a8111ca69d42c71632c3f39d23f67965d Mon Sep 17 00:00:00 2001 From: AIRONAX Developer Date: Fri, 2 Oct 2026 20:29:58 +0530 Subject: [PATCH 02/14] refactor(web): consolidate flash, API key, and rate-limit handlers Fewer web entrypoints; shared flash cookies and unified rate limiting without behavior regressions on public routes. --- core/server/web/apikey_flash.go | 22 ------ .../{apikey_create.go => apikey_handlers.go} | 38 ++++++++-- core/server/web/apikey_reveal.go | 36 ---------- core/server/web/apps_flash.go | 2 +- ...error_flash_cookie.go => flash_cookies.go} | 45 +++++++++++- core/server/web/page_flash.go | 30 -------- core/server/web/portal_auth.go | 45 ------------ core/server/web/portal_handlers.go | 38 ++++++++++ core/server/web/rate_limit.go | 70 +++++++++++++++++- core/server/web/ratelimit.go | 71 ------------------- core/server/web/settings_field_acl.go | 2 +- core/server/web/settings_model_acl.go | 2 +- core/server/web/swc_workspace.go | 1 + .../server/web/query_flash_sanitize_test.go | 28 ++++++++ test/core/server/web/rate_limit_test.go | 24 +++++++ 15 files changed, 238 insertions(+), 216 deletions(-) delete mode 100644 core/server/web/apikey_flash.go rename core/server/web/{apikey_create.go => apikey_handlers.go} (53%) delete mode 100644 core/server/web/apikey_reveal.go rename core/server/web/{record_error_flash_cookie.go => flash_cookies.go} (50%) delete mode 100644 core/server/web/page_flash.go delete mode 100644 core/server/web/portal_auth.go delete mode 100644 core/server/web/ratelimit.go create mode 100644 test/core/server/web/query_flash_sanitize_test.go create mode 100644 test/core/server/web/rate_limit_test.go diff --git a/core/server/web/apikey_flash.go b/core/server/web/apikey_flash.go deleted file mode 100644 index 4335fe4..0000000 --- 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 104b1b6..f5766cb 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 39e7274..0000000 --- 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 88b3df6..952b7d3 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/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 396d32a..c95fadd 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/page_flash.go b/core/server/web/page_flash.go deleted file mode 100644 index 20d23b7..0000000 --- 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 ea77c75..0000000 --- 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 51df60d..3d7db81 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/rate_limit.go b/core/server/web/rate_limit.go index 0cb476c..18fc211 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 d50c6dd..0000000 --- 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/settings_field_acl.go b/core/server/web/settings_field_acl.go index c5be050..36cb503 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 e9dbe53..a468dce 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/swc_workspace.go b/core/server/web/swc_workspace.go index a133ac2..5cda26e 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/test/core/server/web/query_flash_sanitize_test.go b/test/core/server/web/query_flash_sanitize_test.go new file mode 100644 index 0000000..75401cc --- /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/rate_limit_test.go b/test/core/server/web/rate_limit_test.go new file mode 100644 index 0000000..0c1120c --- /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") + } +} From 24add05f9869772125e8c4990756997220a1a88e Mon Sep 17 00:00:00 2001 From: AIRONAX Developer Date: Fri, 2 Oct 2026 20:30:07 +0530 Subject: [PATCH 03/14] feat(bus): SUM-PLAT-15 durable events, NOTIFY fan-out, and channel ACL Persist bus events, publish on ORM hooks, and harden SWC listen/reconnect with channel-scoped delivery. --- core/orm/bus_notify.go | 18 + core/orm/bus_publish_hook.go | 28 ++ core/orm/sys_bus_event.go | 136 ++++++++ core/server/run.go | 1 + core/server/web/swc_bus.go | 325 +++++++++++++++++++ core/server/web/swc_bus_channel.go | 79 +++++ core/server/web/swc_bus_hub.go | 149 --------- core/server/web/swc_bus_publish.go | 78 +++++ core/swc/src/services/bus.ts | 132 +++++++- test/core/server/web/swc_bus_channel_test.go | 25 ++ test/core/server/web/swc_bus_hub_test.go | 32 +- 11 files changed, 824 insertions(+), 179 deletions(-) create mode 100644 core/orm/bus_notify.go create mode 100644 core/orm/bus_publish_hook.go create mode 100644 core/orm/sys_bus_event.go create mode 100644 core/server/web/swc_bus_channel.go delete mode 100644 core/server/web/swc_bus_hub.go create mode 100644 core/server/web/swc_bus_publish.go create mode 100644 test/core/server/web/swc_bus_channel_test.go diff --git a/core/orm/bus_notify.go b/core/orm/bus_notify.go new file mode 100644 index 0000000..6543921 --- /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 0000000..f350111 --- /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/sys_bus_event.go b/core/orm/sys_bus_event.go new file mode 100644 index 0000000..b7c613f --- /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/server/run.go b/core/server/run.go index 8cc0185..5cb0d23 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/swc_bus.go b/core/server/web/swc_bus.go index 99367ba..aa773ea 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 0000000..cc72407 --- /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 7bba43c..0000000 --- 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 0000000..18bfa5c --- /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/swc/src/services/bus.ts b/core/swc/src/services/bus.ts index 4329e2b..20e72f6 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/test/core/server/web/swc_bus_channel_test.go b/test/core/server/web/swc_bus_channel_test.go new file mode 100644 index 0000000..ae3bfe0 --- /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 cce33bc..8232c0a 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: From 0bb3ac020c9bc54d022e62ea596361ee3999fcce Mon Sep 17 00:00:00 2001 From: AIRONAX Developer Date: Fri, 2 Oct 2026 20:30:35 +0530 Subject: [PATCH 04/14] feat(mail): SUM-PLAT-14 followers, notifications, inbox bell, and mail queue Mail thread on models, follower-driven notifications, SWC bell, and outbound HTML mail via the worker. --- addons/base/models/zmodels.go | 3 + addons/mail/followers.go | 171 +++++++++++++++++++++ addons/mail/hooks.go | 61 ++++++++ addons/mail/mail.go | 10 +- addons/mail/manifest.json | 1 + addons/mail/models/core_user_mail.go | 11 ++ addons/mail/models/mail_follower.go | 16 ++ addons/mail/models/mail_message_subtype.go | 16 ++ addons/mail/models/mail_notification.go | 18 +++ addons/mail/models/zmodels.go | 4 + addons/mail/models/zrefs.go | 3 + addons/mail/security/sys.access.csv | 3 + core/engine/assets/css/sumeru-shell.css | 59 +++++++ core/mail/smtp.go | 55 +++++++ core/mail/worker.go | 21 ++- core/modelmeta/model_spec.go | 15 +- core/modelmeta/tags.go | 3 + core/modelreg/activate.go | 4 + core/modelreg/register.go | 1 + core/orm/mail_thread.go | 26 ++++ core/ormmodels/zmodels.go | 1 + core/server/web/swc_notifications.go | 150 ++++++++++++++++++ core/swc/src/main.ts | 2 +- core/swc/src/model/record.ts | 13 ++ core/swc/src/shell/notification-bell.ts | 80 ++++++++++ core/swc/src/shell/shell-chrome.ts | 8 +- test/core/coverage/exports_test.go | 3 +- 27 files changed, 747 insertions(+), 11 deletions(-) create mode 100644 addons/mail/followers.go create mode 100644 addons/mail/hooks.go create mode 100644 addons/mail/models/core_user_mail.go create mode 100644 addons/mail/models/mail_follower.go create mode 100644 addons/mail/models/mail_message_subtype.go create mode 100644 addons/mail/models/mail_notification.go create mode 100644 core/orm/mail_thread.go create mode 100644 core/server/web/swc_notifications.go create mode 100644 core/swc/src/shell/notification-bell.ts diff --git a/addons/base/models/zmodels.go b/addons/base/models/zmodels.go index 300693b..5150e40 100644 --- a/addons/base/models/zmodels.go +++ b/addons/base/models/zmodels.go @@ -16,10 +16,13 @@ func init() { &CorePartner{}, &CoreUser{}, &CoreUserAPIKey{}, + &CoreUserIdentity{}, &CoreUserLog{}, + &CoreUserTrustedDevice{}, &ResConfigSettings{}, &SysAttachment{}, &SysAudit{}, + &SysAuthProvider{}, &SysBulkImport{}, &SysConfigParameter{}, &SysFieldAccess{}, diff --git a/addons/mail/followers.go b/addons/mail/followers.go new file mode 100644 index 0000000..2b75dca --- /dev/null +++ b/addons/mail/followers.go @@ -0,0 +1,171 @@ +package mail + +import ( + "context" + "fmt" + "html" + "regexp" + "strings" + + pkgmail "sumeru/core/mail" + "sumeru/core/orm" + "time" +) + +var mentionPattern = regexp.MustCompile(`@([a-zA-Z0-9._-]+)`) + +// SubscribeUserFollower ensures user_id follows the record (idempotent). +func SubscribeUserFollower(ctx context.Context, model string, resID int64, userID int) error { + if userID <= 0 || resID <= 0 || strings.TrimSpace(model) == "" { + return nil + } + if _, ok := orm.Registry["mail.follower"]; !ok { + return nil + } + existing, err := orm.SearchOne(ctx, "mail.follower", map[string]interface{}{ + "res_model": model, + "res_id": int(resID), + "user_id": userID, + }) + if err == nil && existing != nil { + return nil + } + inst := orm.Registry["mail.follower"] + _, err = orm.Create(ctx, inst, map[string]interface{}{ + "res_model": model, + "res_id": int(resID), + "user_id": userID, + "active": true, + }) + return err +} + +// NotifyMessageFollowers creates inbox notifications and optional email for a posted message. +func NotifyMessageFollowers(ctx context.Context, messageID int, authorUserID int) error { + if messageID <= 0 || orm.DB == nil { + return nil + } + msg, err := orm.SearchOne(ctx, "mail.message", map[string]interface{}{"id": messageID}) + if err != nil || msg == nil { + return err + } + model := rowString(msg, "model") + resID := int(rowInt64(msg, "core_id")) + subtype := rowString(msg, "subtype") + body := rowString(msg, "body") + if model == "" || resID <= 0 { + return nil + } + recipients := map[int]struct{}{} + followers, _ := orm.Search(ctx, "mail.follower", [][]interface{}{ + {"res_model", "=", model}, + {"res_id", "=", resID}, + {"active", "=", true}, + }) + for _, f := range followers { + uid := int(rowInt64(f, "user_id")) + if uid <= 0 || uid == authorUserID { + continue + } + recipients[uid] = struct{}{} + } + for _, login := range mentionPattern.FindAllStringSubmatch(body, -1) { + if len(login) < 2 { + continue + } + uid := userIDByLogin(ctx, login[1]) + if uid > 0 && uid != authorUserID { + recipients[uid] = struct{}{} + _ = SubscribeUserFollower(ctx, model, int64(resID), uid) + } + } + authorName := rowString(msg, "author") + subject := fmt.Sprintf("[%s] %s", model, truncate(body, 80)) + for uid := range recipients { + if err := orm.CheckModelAccess(ctx, uid, model, "read"); err != nil { + continue + } + notifID, err := createNotification(ctx, uid, messageID, model, resID) + if err != nil { + continue + } + orm.PublishBusChannel(ctx, orm.UserNotificationsBusChannel(uid), map[string]interface{}{ + "notification_id": notifID, + "message_id": messageID, + "model": model, + "id": resID, + }) + if shouldEmailUser(ctx, uid) { + email := userEmail(ctx, uid) + if email != "" { + htmlBody := renderNotificationHTML(authorName, body, model, resID) + pkgmail.EnqueueHTML(ctx, email, subject, htmlBody, stripHTML(body)) + } + } + } + _ = subtype + return nil +} + +func createNotification(ctx context.Context, uid, messageID int, model string, resID int) (int, error) { + inst := orm.Registry["mail.notification"] + id, err := orm.Create(ctx, inst, map[string]interface{}{ + "user_id": uid, + "message_id": messageID, + "is_read": false, + "record_model": model, + "record_id": resID, + "create_date": time.Now().UTC().Format(time.RFC3339), + }) + return id, err +} + +func userIDByLogin(ctx context.Context, login string) int { + login = strings.TrimSpace(login) + if login == "" || orm.DB == nil { + return 0 + } + rec, err := orm.SearchOne(ctx, "core.user", map[string]interface{}{"login": login}) + if err != nil || rec == nil { + return 0 + } + return int(rowInt64(rec, "id")) +} + +func userEmail(ctx context.Context, uid int) string { + rec, err := orm.SearchOne(ctx, "core.user", map[string]interface{}{"id": uid}) + if err != nil || rec == nil { + return "" + } + return rowString(rec, "email") +} + +func shouldEmailUser(ctx context.Context, uid int) bool { + rec, err := orm.SearchOne(ctx, "core.user", map[string]interface{}{"id": uid}) + if err != nil || rec == nil { + return false + } + if v, ok := rec["notify_email"].(bool); ok { + return v + } + return true +} + +func renderNotificationHTML(author, body, model string, resID int) string { + esc := html.EscapeString(body) + authorEsc := html.EscapeString(author) + return "

" + authorEsc + " on " + html.EscapeString(model) + + " #" + fmt.Sprint(resID) + ":

" + strings.ReplaceAll(esc, "\n", "
") + "

" +} + +func stripHTML(s string) string { + return strings.TrimSpace(s) +} + +func truncate(s string, n int) string { + s = strings.TrimSpace(s) + if len(s) <= n { + return s + } + return s[:n] + "…" +} diff --git a/addons/mail/hooks.go b/addons/mail/hooks.go new file mode 100644 index 0000000..161d908 --- /dev/null +++ b/addons/mail/hooks.go @@ -0,0 +1,61 @@ +package mail + +import ( + "context" + + "sumeru/core/event" + "sumeru/core/orm" +) + +func init() { + event.Subscribe(eventRecordCreated, onMailThreadRecordCreated) + event.Subscribe(eventRecordUpdated, onMailThreadRecordUpdated) +} + +const ( + eventRecordCreated = "record.created" + eventRecordUpdated = "record.updated" +) + +func onMailThreadRecordCreated(ctx context.Context, ev event.Event) error { + model, _ := ev.Payload["model"].(string) + id, _ := payloadInt(ev.Payload["id"]) + if model == "" || id <= 0 || !orm.ModelHasMailThread(model) { + return nil + } + uid := ev.Actor + if uid <= 0 { + uid = orm.SecurityUID(ctx) + } + return SubscribeUserFollower(ctx, model, int64(id), uid) +} + +func onMailThreadRecordUpdated(ctx context.Context, ev event.Event) error { + model, _ := ev.Payload["model"].(string) + id, _ := payloadInt(ev.Payload["id"]) + if model == "" || id <= 0 || !orm.ModelHasMailThread(model) { + return nil + } + rec, err := orm.SearchOne(ctx, model, map[string]interface{}{"id": id}) + if err != nil || rec == nil { + return nil + } + assignee, ok := payloadInt(rec["user_id"]) + if !ok || assignee <= 0 { + return nil + } + return SubscribeUserFollower(ctx, model, int64(id), assignee) +} + +func payloadInt(v interface{}) (int, bool) { + switch t := v.(type) { + case int: + return t, true + case int64: + return int(t), true + case float64: + return int(t), true + default: + return 0, false + } +} diff --git a/addons/mail/mail.go b/addons/mail/mail.go index a83b8a9..596d6b2 100644 --- a/addons/mail/mail.go +++ b/addons/mail/mail.go @@ -123,8 +123,14 @@ func PostMessage(ctx context.Context, model string, coreID int64, body, subtype, if settings, ok := firstCompanyMailSettings(ctx); ok && settings.id > 0 { vals["company_id"] = int(settings.id) } - _, err := orm.Create(ctx, inst, vals) - return err + msgID, err := orm.Create(ctx, inst, vals) + if err != nil { + return err + } + if subtype == SubtypeComment { + _ = NotifyMessageFollowers(ctx, msgID, uid) + } + return nil } // ListCommentsForRecord returns user chatter lines (subtype comment) for a record, oldest first. diff --git a/addons/mail/manifest.json b/addons/mail/manifest.json index a613ef6..e91f82e 100644 --- a/addons/mail/manifest.json +++ b/addons/mail/manifest.json @@ -10,6 +10,7 @@ "security/security.xml", "security/sys.access.csv", "data/mail_activity_type_data.xml", + "data/mail_subtype_data.xml", "views/actions.xml", "views/mail_message_list_views.xml", "views/mail_message_form_views.xml", diff --git a/addons/mail/models/core_user_mail.go b/addons/mail/models/core_user_mail.go new file mode 100644 index 0000000..5e728ea --- /dev/null +++ b/addons/mail/models/core_user_mail.go @@ -0,0 +1,11 @@ +package models + +import ( + "sumeru/core/sdk" +) + +type CoreUserMail struct { + sdk.Model `sumeru:"inherit=core.user"` + + NotifyEmail sdk.Boolean `sumeru:"string=Email Notifications,default=true,column=notify_email"` +} diff --git a/addons/mail/models/mail_follower.go b/addons/mail/models/mail_follower.go new file mode 100644 index 0000000..52fd40e --- /dev/null +++ b/addons/mail/models/mail_follower.go @@ -0,0 +1,16 @@ +package models + +import ( + "sumeru/core/sdk" +) + +type MailFollower struct { + sdk.Model `sumeru:"model=mail.follower"` + + ResModel sdk.String `sumeru:"required,index,column=res_model,string=Document Model"` + ResID sdk.Integer `sumeru:"required,index,column=res_id,string=Document ID"` + UserID sdk.Many2One[CoreUser] `sumeru:"required,index,string=User"` + PartnerID sdk.Many2One[CorePartner] `sumeru:"index,string=Partner"` + SubtypeIds sdk.Many2Many[MailMessageSubtype] `sumeru:"string=Subtypes,table=mail_follower_subtype_rel,left=follower_id,right=subtype_id"` + Active sdk.Boolean `sumeru:"string=Active,default=true"` +} diff --git a/addons/mail/models/mail_message_subtype.go b/addons/mail/models/mail_message_subtype.go new file mode 100644 index 0000000..f06cecc --- /dev/null +++ b/addons/mail/models/mail_message_subtype.go @@ -0,0 +1,16 @@ +package models + +import ( + "sumeru/core/sdk" +) + +type MailMessageSubtype struct { + sdk.Model `sumeru:"model=mail.message.subtype"` + + Name sdk.String `sumeru:"required,unique,string=Technical Name"` + Description sdk.String `sumeru:"string=Description"` + Internal sdk.Boolean `sumeru:"string=Internal,default=false"` + Default sdk.Boolean `sumeru:"string=Default,default=false"` + Hidden sdk.Boolean `sumeru:"string=Hidden,default=false"` + Sequence sdk.Integer `sumeru:"string=Sequence,default=10"` +} diff --git a/addons/mail/models/mail_notification.go b/addons/mail/models/mail_notification.go new file mode 100644 index 0000000..d43fd79 --- /dev/null +++ b/addons/mail/models/mail_notification.go @@ -0,0 +1,18 @@ +package models + +import ( + "sumeru/core/sdk" +) + +type MailNotification struct { + sdk.Model `sumeru:"model=mail.notification"` + + UserID sdk.Many2One[CoreUser] `sumeru:"required,index,string=User"` + MessageID sdk.Many2One[MailMessage] `sumeru:"required,index,string=Message"` + IsRead sdk.Boolean `sumeru:"string=Read,default=false"` + ReadDate sdk.DateTime `sumeru:"string=Read Date"` + RecordModel sdk.String `sumeru:"index,column=record_model,string=Record Model"` + RecordID sdk.Integer `sumeru:"index,column=record_id,string=Record ID"` + CompanyID sdk.Many2One[CoreCompany] `sumeru:"index,string=Company"` + CreateDate sdk.DateTime `sumeru:"required,string=Created"` +} diff --git a/addons/mail/models/zmodels.go b/addons/mail/models/zmodels.go index 08e37f6..d9156c7 100644 --- a/addons/mail/models/zmodels.go +++ b/addons/mail/models/zmodels.go @@ -6,11 +6,15 @@ import "sumeru/core/sdk" func init() { sdk.MustRegister("mail", + &CoreUserMail{}, &MailActivity{}, &MailActivityPlan{}, &MailActivityPlanTemplate{}, &MailActivityType{}, + &MailFollower{}, &MailMessage{}, + &MailMessageSubtype{}, + &MailNotification{}, &MailTemplate{}, ) } diff --git a/addons/mail/models/zrefs.go b/addons/mail/models/zrefs.go index bcec8b9..9c74616 100644 --- a/addons/mail/models/zrefs.go +++ b/addons/mail/models/zrefs.go @@ -9,6 +9,9 @@ import ( // CoreCompany → core.company type CoreCompany = basemodels.CoreCompany +// CorePartner → core.partner +type CorePartner = basemodels.CorePartner + // CoreUser → core.user type CoreUser = basemodels.CoreUser diff --git a/addons/mail/security/sys.access.csv b/addons/mail/security/sys.access.csv index b3a7d2d..18727f7 100644 --- a/addons/mail/security/sys.access.csv +++ b/addons/mail/security/sys.access.csv @@ -5,3 +5,6 @@ access_mail_activity_type_user,access_mail_activity_type_user,mail.activity.type access_mail_activity_plan_user,access_mail_activity_plan_user,mail.activity.plan,base.group_system,1,1,1,1 access_mail_activity_plan_template_user,access_mail_activity_plan_template_user,mail.activity.plan.template,base.group_system,1,1,1,1 access_mail_template_user,access_mail_template_user,mail.template,base.group_system,1,1,1,1 +access_mail_follower_user,access_mail_follower_user,mail.follower,base.group_user,1,1,1,1 +access_mail_subtype_user,access_mail_subtype_user,mail.message.subtype,base.group_user,1,0,0,0 +access_mail_notification_user,access_mail_notification_user,mail.notification,base.group_user,1,1,1,0 diff --git a/core/engine/assets/css/sumeru-shell.css b/core/engine/assets/css/sumeru-shell.css index b0bc731..9f57699 100644 --- a/core/engine/assets/css/sumeru-shell.css +++ b/core/engine/assets/css/sumeru-shell.css @@ -149,6 +149,65 @@ flex-shrink: 0; } +.sum-notification-bell { + position: relative; +} + +.sum-notification-badge { + position: absolute; + top: -4px; + right: -4px; + min-width: 1rem; + padding: 0 0.25rem; + font-size: 0.65rem; + font-weight: 700; + line-height: 1rem; + text-align: center; + border-radius: 999px; + background: var(--sum-danger, #dc2626); + color: #fff; +} + +.sum-notification-panel { + position: absolute; + right: 0.75rem; + top: 3rem; + z-index: 200; + width: min(22rem, 92vw); + max-height: 20rem; + overflow: auto; + background: var(--sum-surface, #fff); + border: 1px solid var(--sum-border, #e5e5e5); + border-radius: 8px; + box-shadow: 0 8px 24px rgb(0 0 0 / 12%); + padding: 0.35rem; +} + +.sum-notification-row { + display: block; + width: 100%; + text-align: left; + padding: 0.5rem 0.65rem; + border: none; + background: transparent; + cursor: pointer; + font: inherit; +} + +.sum-notification-row--unread { + font-weight: 600; +} + +.sum-notification-mark-all { + width: 100%; + margin-top: 0.35rem; + padding: 0.35rem; + border: none; + background: var(--sum-muted-bg, #f5f5f5); + cursor: pointer; + font: inherit; +} + .sum-brand-lockup { display: flex; align-items: center; diff --git a/core/mail/smtp.go b/core/mail/smtp.go index eca52a4..e7afedb 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 bff9df1..28ed0ac 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 7c46276..35d70c5 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 5f57659..d56cccd 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 8281b89..6cd8a70 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 72cf34e..abfb318 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/mail_thread.go b/core/orm/mail_thread.go new file mode 100644 index 0000000..1cf5fc4 --- /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/ormmodels/zmodels.go b/core/ormmodels/zmodels.go index 339b83a..5290d57 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/server/web/swc_notifications.go b/core/server/web/swc_notifications.go new file mode 100644 index 0000000..9fdd807 --- /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/swc/src/main.ts b/core/swc/src/main.ts index e4a2d36..927a4d7 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 07eb27c..8d5d55f 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/shell/notification-bell.ts b/core/swc/src/shell/notification-bell.ts new file mode 100644 index 0000000..42548d6 --- /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 5f82e14..c8ed5c9 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/test/core/coverage/exports_test.go b/test/core/coverage/exports_test.go index 5bcaeef..7b72ee2 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" { From cb086323eda056450b65883fcd7e3aad0009d9b0 Mon Sep 17 00:00:00 2001 From: AIRONAX Developer Date: Fri, 2 Oct 2026 20:30:42 +0530 Subject: [PATCH 05/14] feat(auth): SUM-PLAT-13 OIDC PKCE, JWKS id_token verify, and login MFA step Provider and identity models, OAuth callback with fail-closed JWT checks, TOTP/trusted-device cookies, and shared login completion path. --- addons/base/models/core_user_identity.go | 13 + addons/base/models/core_user_trusteddevice.go | 14 + addons/base/models/sys_auth_provider.go | 22 + core/engine/templates/totp_login.html | 49 ++ core/orm/user_totp.go | 146 ++++++ core/security/fields.go | 6 + core/server/auth/oidc_verify.go | 320 ++++++++++++ core/server/auth/totp.go | 52 ++ core/server/auth/totp_qr.go | 24 + core/server/auth/totp_qr_png.go | 11 + core/server/web/cookie_helpers.go | 7 + core/server/web/login.go | 467 +++++++++++++++--- core/server/web/login_csrf.go | 50 -- core/server/web/login_lockout.go | 84 ---- core/server/web/login_next.go | 45 -- core/server/web/oauth_login.go | 406 +++++++++++++++ core/server/web/query_flash.go | 50 +- core/server/web/routes_table.go | 2 + core/server/web/signed_cookie.go | 73 +++ core/server/web/testexports.go | 54 +- core/server/web/web_constants.go | 24 +- go.mod | 2 + go.sum | 2 + test/core/orm/sqlmock_test.go | 14 + test/core/orm/user_totp_test.go | 61 +++ test/core/server/auth/oidc_verify_test.go | 199 ++++++++ test/core/server/auth/totp_test.go | 26 + test/core/server/web/oauth_callback_test.go | 46 ++ .../server/web/pending_mfa_cookie_test.go | 24 + 29 files changed, 2022 insertions(+), 271 deletions(-) create mode 100644 addons/base/models/core_user_identity.go create mode 100644 addons/base/models/core_user_trusteddevice.go create mode 100644 addons/base/models/sys_auth_provider.go create mode 100644 core/engine/templates/totp_login.html create mode 100644 core/orm/user_totp.go create mode 100644 core/server/auth/oidc_verify.go create mode 100644 core/server/auth/totp.go create mode 100644 core/server/auth/totp_qr.go create mode 100644 core/server/auth/totp_qr_png.go delete mode 100644 core/server/web/login_csrf.go delete mode 100644 core/server/web/login_lockout.go delete mode 100644 core/server/web/login_next.go create mode 100644 core/server/web/oauth_login.go create mode 100644 core/server/web/signed_cookie.go create mode 100644 test/core/orm/user_totp_test.go create mode 100644 test/core/server/auth/oidc_verify_test.go create mode 100644 test/core/server/auth/totp_test.go create mode 100644 test/core/server/web/oauth_callback_test.go create mode 100644 test/core/server/web/pending_mfa_cookie_test.go diff --git a/addons/base/models/core_user_identity.go b/addons/base/models/core_user_identity.go new file mode 100644 index 0000000..e3fe60a --- /dev/null +++ b/addons/base/models/core_user_identity.go @@ -0,0 +1,13 @@ +package models + +import ( + "sumeru/core/sdk" +) + +type CoreUserIdentity struct { + sdk.Model `sumeru:"model=core.user.identity"` + + ProviderID sdk.Many2One[SysAuthProvider] `sumeru:"required,index,string=Provider"` + Subject sdk.String `sumeru:"required,index,string=Subject"` + UserID sdk.Many2One[CoreUser] `sumeru:"required,index,string=User"` +} diff --git a/addons/base/models/core_user_trusteddevice.go b/addons/base/models/core_user_trusteddevice.go new file mode 100644 index 0000000..cd18ea6 --- /dev/null +++ b/addons/base/models/core_user_trusteddevice.go @@ -0,0 +1,14 @@ +package models + +import ( + "sumeru/core/sdk" +) + +type CoreUserTrustedDevice struct { + sdk.Model `sumeru:"model=core.user.trusteddevice"` + + UserID sdk.Many2One[CoreUser] `sumeru:"required,index,string=User"` + TokenHash sdk.String `sumeru:"required,index,string=Token Hash,column=token_hash"` + UserAgent sdk.String `sumeru:"string=User Agent,column=user_agent"` + ExpiresAt sdk.DateTime `sumeru:"required,index,string=Expires At,column=expires_at"` +} diff --git a/addons/base/models/sys_auth_provider.go b/addons/base/models/sys_auth_provider.go new file mode 100644 index 0000000..6871708 --- /dev/null +++ b/addons/base/models/sys_auth_provider.go @@ -0,0 +1,22 @@ +package models + +import ( + "sumeru/core/sdk" +) + +type SysAuthProvider struct { + sdk.Model `sumeru:"model=sys.auth.provider"` + + Name sdk.String `sumeru:"required,string=Name"` + ProviderType sdk.String `sumeru:"required,string=Type,default=oidc,selection=oidc:OpenID Connect,oauth2:OAuth2"` + ClientID sdk.String `sumeru:"required,string=Client ID"` + ClientSecret sdk.String `sumeru:"string=Client Secret"` + IssuerURL sdk.String `sumeru:"string=Issuer URL"` + AuthorizeURL sdk.String `sumeru:"string=Authorize URL"` + TokenURL sdk.String `sumeru:"string=Token URL"` + JwksURL sdk.String `sumeru:"string=JWKS URL"` + Scopes sdk.String `sumeru:"string=Scopes,default=openid email profile"` + Enabled sdk.Boolean `sumeru:"string=Enabled,default=false"` + ButtonLabel sdk.String `sumeru:"string=Button Label"` + LinkPolicy sdk.String `sumeru:"string=Unknown Users,default=deny,selection=deny:Deny,link:Link by email,jit_internal:JIT internal,jit_portal:JIT portal"` +} diff --git a/core/engine/templates/totp_login.html b/core/engine/templates/totp_login.html new file mode 100644 index 0000000..86b23bd --- /dev/null +++ b/core/engine/templates/totp_login.html @@ -0,0 +1,49 @@ + + + + + + Two-factor authentication — {{.AppName}} + {{range .Stylesheets}} + + {{end}} + + + + + + + + diff --git a/core/orm/user_totp.go b/core/orm/user_totp.go new file mode 100644 index 0000000..732b862 --- /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/security/fields.go b/core/security/fields.go index 928db55..17763f4 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 0000000..675246d --- /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 0000000..f4101e0 --- /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 0000000..bb7ac8d --- /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 0000000..bd564a1 --- /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/web/cookie_helpers.go b/core/server/web/cookie_helpers.go index 5fb9095..0996453 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/login.go b/core/server/web/login.go index 218a9ba..d7d43d7 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 617099a..0000000 --- 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 98ca07e..0000000 --- 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 309e4ed..0000000 --- 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 0000000..908bda4 --- /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/query_flash.go b/core/server/web/query_flash.go index 5be979b..8e59a58 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/routes_table.go b/core/server/web/routes_table.go index 3cc59e0..7ad6b13 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/signed_cookie.go b/core/server/web/signed_cookie.go new file mode 100644 index 0000000..91bc24f --- /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/testexports.go b/core/server/web/testexports.go index 08f347b..fe8253f 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 9486b7b..5ebaae5 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/go.mod b/go.mod index 5329084..e89e88c 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 1c0d42b..15e5cbe 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/orm/sqlmock_test.go b/test/core/orm/sqlmock_test.go index bfeb44e..fa231b8 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 0000000..75b9959 --- /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 0000000..4ddea3a --- /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 0000000..217496f --- /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/oauth_callback_test.go b/test/core/server/web/oauth_callback_test.go new file mode 100644 index 0000000..fa9467f --- /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 0000000..979407f --- /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) + } +} From 5625a5e956094f9cf3837208280caedf43b8e26b Mon Sep 17 00:00:00 2001 From: AIRONAX Developer Date: Fri, 2 Oct 2026 20:30:50 +0530 Subject: [PATCH 06/14] feat(auth): SUM-PLAT-13 provider admin UI and self-service TOTP enrollment Settings menus for sys.auth.provider and linked identities; account security flow with QR confirm and admin read-only 2FA flag. --- addons/base/manifest.json | 4 + addons/base/views/auth_actions.xml | 7 ++ .../views/core_user_identity_list_views.xml | 10 ++ addons/base/views/core_users_form_views.xml | 2 +- addons/base/views/menus.xml | 2 + .../views/sys_auth_provider_form_views.xml | 41 +++++++ .../views/sys_auth_provider_list_views.xml | 12 ++ .../templates/settings_account_inner.html | 64 ++++++++--- core/server/web/settings_account.go | 15 ++- core/server/web/settings_account_totp.go | 104 ++++++++++++++++++ 10 files changed, 245 insertions(+), 16 deletions(-) create mode 100644 addons/base/views/auth_actions.xml create mode 100644 addons/base/views/core_user_identity_list_views.xml create mode 100644 addons/base/views/sys_auth_provider_form_views.xml create mode 100644 addons/base/views/sys_auth_provider_list_views.xml create mode 100644 core/server/web/settings_account_totp.go diff --git a/addons/base/manifest.json b/addons/base/manifest.json index 70e9554..27155d5 100644 --- a/addons/base/manifest.json +++ b/addons/base/manifest.json @@ -21,6 +21,10 @@ "views/sys_report_action_list_views.xml", "views/sys_report_action_form_views.xml", "views/security_actions.xml", + "views/auth_actions.xml", + "views/sys_auth_provider_list_views.xml", + "views/sys_auth_provider_form_views.xml", + "views/core_user_identity_list_views.xml", "views/core_company_form_views.xml", "views/core_company_kanban_views.xml", "views/core_company_list_views.xml", diff --git a/addons/base/views/auth_actions.xml b/addons/base/views/auth_actions.xml new file mode 100644 index 0000000..a06e99d --- /dev/null +++ b/addons/base/views/auth_actions.xml @@ -0,0 +1,7 @@ + + + + + + + diff --git a/addons/base/views/core_user_identity_list_views.xml b/addons/base/views/core_user_identity_list_views.xml new file mode 100644 index 0000000..a5f3bd6 --- /dev/null +++ b/addons/base/views/core_user_identity_list_views.xml @@ -0,0 +1,10 @@ + + + + + + + + + + diff --git a/addons/base/views/core_users_form_views.xml b/addons/base/views/core_users_form_views.xml index 70a3670..c9c9c92 100644 --- a/addons/base/views/core_users_form_views.xml +++ b/addons/base/views/core_users_form_views.xml @@ -55,7 +55,7 @@ diff --git a/addons/base/views/menus.xml b/addons/base/views/menus.xml index 9349c2b..04e59d2 100644 --- a/addons/base/views/menus.xml +++ b/addons/base/views/menus.xml @@ -20,6 +20,8 @@ + + diff --git a/addons/base/views/sys_auth_provider_form_views.xml b/addons/base/views/sys_auth_provider_form_views.xml new file mode 100644 index 0000000..15ad985 --- /dev/null +++ b/addons/base/views/sys_auth_provider_form_views.xml @@ -0,0 +1,41 @@ + + + + + +
+

+ +

+
+ + + + + + + + + + + + + + + + + + + + + + + + + +
+
+
+
diff --git a/addons/base/views/sys_auth_provider_list_views.xml b/addons/base/views/sys_auth_provider_list_views.xml new file mode 100644 index 0000000..10d451f --- /dev/null +++ b/addons/base/views/sys_auth_provider_list_views.xml @@ -0,0 +1,12 @@ + + + + + + + + + + + + diff --git a/core/engine/templates/settings_account_inner.html b/core/engine/templates/settings_account_inner.html index c4ef597..f60a89c 100644 --- a/core/engine/templates/settings_account_inner.html +++ b/core/engine/templates/settings_account_inner.html @@ -2,23 +2,61 @@

Account security

-

Change your login password. You will stay signed in on this browser.

+

Manage your password and two-factor authentication.

Back to Settings
- diff --git a/core/server/web/settings_account.go b/core/server/web/settings_account.go index 4093029..2199eb5 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 0000000..ae804dd --- /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 +} From f110961308bf3786d3877c791d6d942aca62be0d Mon Sep 17 00:00:00 2001 From: AIRONAX Developer Date: Fri, 2 Oct 2026 20:30:57 +0530 Subject: [PATCH 07/14] feat(web): SUM-PLAT-13 enterprise login and setup shell with SSO buttons Split brand/sign-in layout, enabled IdP buttons, auth.local_enabled behavior, and aligned setup wizard styling. --- core/engine/assets/css/sumeru-login.css | 237 ++++++++++++++++++++++- core/engine/templates/login.html | 100 +++++++--- core/engine/templates/setup.html | 16 +- test/core/server/web/login_csrf_test.go | 35 ++++ test/core/server/web/query_flash_test.go | 17 ++ 5 files changed, 372 insertions(+), 33 deletions(-) diff --git a/core/engine/assets/css/sumeru-login.css b/core/engine/assets/css/sumeru-login.css index 6701e89..f03fffc 100644 --- a/core/engine/assets/css/sumeru-login.css +++ b/core/engine/assets/css/sumeru-login.css @@ -3,10 +3,94 @@ .sum-login-page { background-color: var(--sum-bg); min-height: 100vh; + font-family: var(--sum-font-sans); + margin: 0; +} + +.sum-login-shell { + min-height: 100vh; + display: flex; + flex-direction: row; + align-items: stretch; +} + +.sum-login-brand { + flex: 1 1 42%; + background: linear-gradient(145deg, var(--sum-header) 0%, var(--sum-header-dark) 55%, #1a2332 100%); + color: var(--sum-surface); display: flex; align-items: center; justify-content: center; - font-family: var(--sum-font-sans); + padding: 3rem 2.5rem; +} + +.sum-login-brand-inner { + max-width: 420px; +} + +.sum-login-brand-logo { + margin: 0 0 1.75rem; + width: 72px; + height: 72px; +} + +.sum-login-brand-kicker { + margin: 0 0 0.5rem; + font-size: var(--sum-font-size-sm); + font-weight: var(--sum-font-weight-semibold); + letter-spacing: 0.08em; + text-transform: uppercase; + opacity: 0.85; +} + +.sum-login-brand-title { + margin: 0 0 1rem; + font-size: 1.75rem; + font-weight: var(--sum-font-weight-bold); + line-height: 1.2; + letter-spacing: -0.02em; +} + +.sum-login-brand-copy { + margin: 0 0 1.5rem; + font-size: var(--sum-font-size-md); + line-height: 1.55; + opacity: 0.92; +} + +.sum-login-trust-list { + margin: 0; + padding: 0; + list-style: none; + font-size: var(--sum-font-size-sm); + line-height: 1.6; + opacity: 0.9; +} + +.sum-login-trust-list li { + position: relative; + padding-left: 1.25rem; + margin-bottom: 0.35rem; +} + +.sum-login-trust-list li::before { + content: ""; + position: absolute; + left: 0; + top: 0.55em; + width: 6px; + height: 6px; + border-radius: 50%; + background: var(--sum-surface); + opacity: 0.75; +} + +.sum-login-panel { + flex: 1 1 58%; + display: flex; + align-items: center; + justify-content: center; + padding: 2rem 1.5rem; } .sum-login-card { @@ -19,6 +103,137 @@ border: 1px solid var(--sum-line); } +.sum-login-card--panel { + max-width: 440px; +} + +.sum-login-sso-list { + display: flex; + flex-direction: column; + gap: 0.65rem; + margin-bottom: 0.5rem; +} + +.sum-login-sso { + display: block; + text-align: center; + padding: 0.75rem 1rem; + border-radius: var(--sum-radius-sm); + border: 1px solid var(--sum-line-strong); + background: var(--sum-surface); + color: var(--sum-header); + font-weight: var(--sum-font-weight-semibold); + font-size: var(--sum-font-size-md); + text-decoration: none; + transition: background var(--sum-dur-slow) var(--sum-ease), border-color var(--sum-dur-slow) var(--sum-ease); +} + +.sum-login-sso:hover { + background: var(--sum-gray-50); + border-color: var(--sum-header); +} + +.sum-login-divider { + display: flex; + align-items: center; + gap: 1rem; + margin: 1.5rem 0; + color: var(--sum-muted); + font-size: var(--sum-font-size-sm); + text-transform: uppercase; + letter-spacing: var(--sum-letter-spacing-ui); +} + +.sum-login-divider::before, +.sum-login-divider::after { + content: ""; + flex: 1; + height: 1px; + background: var(--sum-line); +} + +.sum-login-divider span { + flex-shrink: 0; +} + +.sum-login-info { + background: var(--sum-gray-50); + border: 1px solid var(--sum-line); + color: var(--sum-muted-ink); + padding: 0.75rem 1rem; + border-radius: var(--sum-radius-sm); + font-size: var(--sum-font-size-sm); + margin-bottom: 1.25rem; + line-height: 1.45; +} + +.sum-login-muted { + color: var(--sum-muted); + font-size: var(--sum-font-size-md); + text-align: center; + line-height: 1.5; +} + +.sum-login-footer { + margin-top: 2rem; + padding-top: 1.25rem; + border-top: 1px solid var(--sum-line); + font-size: var(--sum-font-size-sm); + color: var(--sum-muted); + text-align: center; +} + +.sum-login-footer-sep { + margin: 0 0.35rem; +} + +.sum-login-trust-device { + display: flex; + align-items: center; + gap: 0.5rem; + font-size: var(--sum-font-size-sm); + color: var(--sum-muted-ink); + margin: 0.5rem 0 0; + cursor: pointer; +} + +.sum-login-totp-input { + font-size: 1.35rem !important; + letter-spacing: 0.35em; + text-align: center; + font-variant-numeric: tabular-nums; +} + +.sum-login-back { + margin: 1.25rem 0 0; + text-align: center; + font-size: var(--sum-font-size-sm); +} + +.sum-login-back a { + color: var(--sum-header); + font-weight: var(--sum-font-weight-medium); +} + +@media (max-width: 900px) { + .sum-login-shell { + flex-direction: column; + } + + .sum-login-brand { + flex: none; + padding: 2rem 1.5rem; + } + + .sum-login-brand-title { + font-size: 1.35rem; + } + + .sum-login-trust-list { + display: none; + } +} + .sum-login-card .sum-login-logo { display: block; margin: 0 auto 2rem; @@ -244,6 +459,26 @@ max-width: 520px; } +.sum-settings-account-section { + margin-top: 2rem; + padding-top: 1.5rem; + border-top: 1px solid var(--sum-line); +} + +.sum-settings-section-title { + margin: 0 0 0.75rem; + font-size: var(--sum-font-size-md); + font-weight: var(--sum-font-weight-bold); + color: var(--sum-header); +} + +.sum-settings-totp-qr { + display: block; + margin: 0 0 1.25rem; + border: 1px solid var(--sum-line); + border-radius: var(--sum-radius-sm); +} + .sum-setup-progress { margin: 0 0 0.35rem; font-size: var(--sum-font-size-sm); diff --git a/core/engine/templates/login.html b/core/engine/templates/login.html index 079a7ca..fb49d7b 100644 --- a/core/engine/templates/login.html +++ b/core/engine/templates/login.html @@ -3,7 +3,7 @@ - Sign in — Sumeru + Sign in — {{.AppName}} {{range .Stylesheets}} {{end}} @@ -12,36 +12,78 @@ -