mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-05 07:26:36 +08:00
refactor(repo): consolidate openflare-server to root and move subprojects to internal/apps
- Merge all files inside openflare-server to the repository root directory. - Relocate agent, relay, and flared subprojects from internal/ to internal/apps/. - Combine docker-compose files and update build context paths to root. - Update GitHub workflows and Dockerfiles to refer to new directories and package names. - Rewrite Go package imports across all files. - Resolve database renew test race condition and clean up docs.
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package analytics provides ClickHouse data access for analytics tables.
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// CountAccessLogs returns the number of access logs matching filter.
|
||||
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
var count int64
|
||||
query := applyFilter(ch.Model(&analyticsmodel.UserAccessLog{}), filter)
|
||||
if err := query.Count(&count).Error; err != nil {
|
||||
return 0, fmt.Errorf("count access logs: %w", err)
|
||||
}
|
||||
return safeUint64Count(count), nil
|
||||
}
|
||||
|
||||
// ListAccessLogs returns paginated access logs and the total match count.
|
||||
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]analyticsmodel.UserAccessLog, uint64, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
|
||||
return []analyticsmodel.UserAccessLog{}, 0, nil
|
||||
}
|
||||
|
||||
var total int64
|
||||
baseQuery := applyFilter(ch.Model(&analyticsmodel.UserAccessLog{}), filter)
|
||||
if err := baseQuery.Count(&total).Error; err != nil {
|
||||
return nil, 0, fmt.Errorf("count access logs: %w", err)
|
||||
}
|
||||
if total == 0 {
|
||||
return []analyticsmodel.UserAccessLog{}, 0, nil
|
||||
}
|
||||
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 {
|
||||
pageSize = 20
|
||||
}
|
||||
offset := (page - 1) * pageSize
|
||||
|
||||
var logs []analyticsmodel.UserAccessLog
|
||||
err := applyFilter(ch.Model(&analyticsmodel.UserAccessLog{}), filter).
|
||||
Order("created_at DESC, id DESC").
|
||||
Limit(pageSize).
|
||||
Offset(offset).
|
||||
Find(&logs).Error
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("list access logs: %w", err)
|
||||
}
|
||||
|
||||
return logs, safeUint64Count(total), nil
|
||||
}
|
||||
|
||||
func safeUint64Count(count int64) uint64 {
|
||||
if count < 0 {
|
||||
return 0
|
||||
}
|
||||
return uint64(count)
|
||||
}
|
||||
|
||||
func applyFilter(query *gorm.DB, filter AccessLogFilter) *gorm.DB {
|
||||
if filter.UserIDs != nil {
|
||||
if len(filter.UserIDs) == 0 {
|
||||
return query.Where("1 = 0")
|
||||
}
|
||||
query = query.Where("user_id IN ?", filter.UserIDs)
|
||||
}
|
||||
if filter.Path != "" {
|
||||
query = query.Where("path LIKE ?", "%"+filter.Path+"%")
|
||||
}
|
||||
if filter.StartTime != nil {
|
||||
query = query.Where("created_at >= ?", *filter.StartTime)
|
||||
}
|
||||
if filter.EndTime != nil {
|
||||
query = query.Where("created_at <= ?", *filter.EndTime)
|
||||
}
|
||||
return query
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import "time"
|
||||
|
||||
// AccessLogFilter scopes ClickHouse user access log queries.
|
||||
type AccessLogFilter struct {
|
||||
// UserIDs filters by user IDs. nil means no user filter; an empty slice means no matches.
|
||||
UserIDs []uint64
|
||||
Path string
|
||||
// StartTime filters created_at >= StartTime when non-nil.
|
||||
StartTime *time.Time
|
||||
// EndTime filters created_at <= EndTime when non-nil.
|
||||
EndTime *time.Time
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||
)
|
||||
|
||||
const hoursInDay = 24
|
||||
|
||||
// DailyTrend is a single day's access count.
|
||||
type DailyTrend struct {
|
||||
Date string
|
||||
Count uint64
|
||||
}
|
||||
|
||||
// BrowserShare is a browser group's share of access logs.
|
||||
type BrowserShare struct {
|
||||
Browser string
|
||||
Count uint64
|
||||
}
|
||||
|
||||
// TopUser is an active user ranked by access count.
|
||||
type TopUser struct {
|
||||
UserID uint64
|
||||
Count uint64
|
||||
}
|
||||
|
||||
// GetDailyTrend returns per-day access counts for the last days days (inclusive of today).
|
||||
func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
|
||||
if days < 1 {
|
||||
days = 7
|
||||
}
|
||||
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
startTime := time.Now().AddDate(0, 0, -(days - 1)).Truncate(hoursInDay * time.Hour)
|
||||
tableName := analyticsmodel.UserAccessLog{}.TableName()
|
||||
|
||||
query := fmt.Sprintf(`
|
||||
SELECT toDate(created_at) AS date, count() AS count
|
||||
FROM %s
|
||||
WHERE created_at >= ?
|
||||
GROUP BY date
|
||||
ORDER BY date ASC
|
||||
`, tableName)
|
||||
|
||||
type trendRow struct {
|
||||
Date time.Time
|
||||
Count uint64
|
||||
}
|
||||
|
||||
var rows []trendRow
|
||||
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
|
||||
return nil, fmt.Errorf("get daily trend: %w", err)
|
||||
}
|
||||
|
||||
trendMap := make(map[string]uint64, days)
|
||||
for i := 0; i < days; i++ {
|
||||
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
|
||||
trendMap[dateStr] = 0
|
||||
}
|
||||
for _, row := range rows {
|
||||
dateStr := row.Date.Format("2006-01-02")
|
||||
trendMap[dateStr] = row.Count
|
||||
}
|
||||
|
||||
result := make([]DailyTrend, 0, days)
|
||||
for i := days - 1; i >= 0; i-- {
|
||||
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
|
||||
result = append(result, DailyTrend{
|
||||
Date: dateStr,
|
||||
Count: trendMap[dateStr],
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetBrowserDistribution returns browser-grouped access counts since startTime.
|
||||
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
tableName := analyticsmodel.UserAccessLog{}.TableName()
|
||||
query := fmt.Sprintf(`
|
||||
SELECT user_agent, count() AS count
|
||||
FROM %s
|
||||
WHERE created_at >= ?
|
||||
GROUP BY user_agent
|
||||
`, tableName)
|
||||
|
||||
type uaRow struct {
|
||||
UserAgent string
|
||||
Count uint64
|
||||
}
|
||||
|
||||
var rows []uaRow
|
||||
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
|
||||
return nil, fmt.Errorf("get browser distribution: %w", err)
|
||||
}
|
||||
|
||||
browserCounts := make(map[string]uint64)
|
||||
for _, row := range rows {
|
||||
browser := ParseBrowserName(row.UserAgent)
|
||||
browserCounts[browser] += row.Count
|
||||
}
|
||||
|
||||
result := make([]BrowserShare, 0, len(browserCounts))
|
||||
for browser, count := range browserCounts {
|
||||
result = append(result, BrowserShare{
|
||||
Browser: browser,
|
||||
Count: count,
|
||||
})
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
return result[i].Count > result[j].Count
|
||||
})
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetTopActiveUsers returns the most active users since startTime.
|
||||
func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
|
||||
if limit < 1 {
|
||||
limit = 10
|
||||
}
|
||||
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
tableName := analyticsmodel.UserAccessLog{}.TableName()
|
||||
query := fmt.Sprintf(`
|
||||
SELECT user_id, count() AS count
|
||||
FROM %s
|
||||
WHERE created_at >= ? AND user_id > 0
|
||||
GROUP BY user_id
|
||||
ORDER BY count DESC
|
||||
LIMIT ?
|
||||
`, tableName)
|
||||
|
||||
var users []TopUser
|
||||
if err := ch.Raw(query, startTime, limit).Scan(&users).Error; err != nil {
|
||||
return nil, fmt.Errorf("get top active users: %w", err)
|
||||
}
|
||||
return users, nil
|
||||
}
|
||||
@@ -0,0 +1,207 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/column"
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupChGormDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
|
||||
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, gormDB.AutoMigrate(&analyticsmodel.UserAccessLog{}))
|
||||
db.SetChDBForTest(gormDB)
|
||||
return gormDB
|
||||
}
|
||||
|
||||
func TestParseBrowserName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ua string
|
||||
want string
|
||||
}{
|
||||
{name: "chrome", ua: "Mozilla/5.0 Chrome/120.0.0.0", want: "Chrome"},
|
||||
{name: "firefox", ua: "Mozilla/5.0 Firefox/121.0", want: "Firefox"},
|
||||
{name: "safari", ua: "Mozilla/5.0 Safari/605.1.15", want: "Safari"},
|
||||
{name: "edge", ua: "Mozilla/5.0 Edg/120.0.0.0", want: "Edge"},
|
||||
{name: "wechat", ua: "MicroMessenger/8.0", want: "WeChat"},
|
||||
{name: "postman", ua: "PostmanRuntime/7.36.0", want: "Postman"},
|
||||
{name: "other", ua: "curl/8.0", want: "Other"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, ParseBrowserName(tt.ua))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
|
||||
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(0), count)
|
||||
}
|
||||
|
||||
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
|
||||
logs, total, err := ListAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}}, 1, 20)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(0), total)
|
||||
assert.Empty(t, logs)
|
||||
}
|
||||
|
||||
func TestListAccessLogs_WithFilters(t *testing.T) {
|
||||
gormDB := setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
|
||||
now := time.Now().UTC().Truncate(time.Second)
|
||||
logs := []analyticsmodel.UserAccessLog{
|
||||
{ID: 1, UserID: 10, Path: "/api/v1/users", Method: "GET", Status: 200, CreatedAt: now},
|
||||
{ID: 2, UserID: 20, Path: "/api/v1/admin/logs", Method: "GET", Status: 200, CreatedAt: now},
|
||||
{ID: 3, UserID: 10, Path: "/api/v1/other", Method: "POST", Status: 201, CreatedAt: now},
|
||||
}
|
||||
require.NoError(t, gormDB.Create(&logs).Error)
|
||||
|
||||
start := now.Add(-time.Hour)
|
||||
filter := AccessLogFilter{
|
||||
UserIDs: []uint64{10},
|
||||
Path: "users",
|
||||
StartTime: &start,
|
||||
}
|
||||
|
||||
count, err := CountAccessLogs(context.Background(), filter)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(1), count)
|
||||
|
||||
result, total, err := ListAccessLogs(context.Background(), filter, 1, 10)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(1), total)
|
||||
require.Len(t, result, 1)
|
||||
assert.Equal(t, uint64(1), result[0].ID)
|
||||
assert.Equal(t, "/api/v1/users", result[0].Path)
|
||||
}
|
||||
|
||||
func TestBatchInsert_Empty(t *testing.T) {
|
||||
err := BatchInsert(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestBatchInsert_UsesModelBatchSQL(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mockBatch := &mockBatch{}
|
||||
mockConn := &mockConn{
|
||||
batch: mockBatch,
|
||||
batchQuery: analyticsmodel.UserAccessLog{}.BatchInsertSQL(),
|
||||
}
|
||||
db.SetChConnForTest(mockConn)
|
||||
t.Cleanup(func() { db.SetChConnForTest(nil) })
|
||||
|
||||
createdAt := time.Now().UTC()
|
||||
err := BatchInsert(ctx, []analyticsmodel.UserAccessLog{
|
||||
{
|
||||
ID: 1,
|
||||
UserID: 42,
|
||||
Path: "/api/v1/test",
|
||||
Method: "GET",
|
||||
IP: "127.0.0.1",
|
||||
UserAgent: "test-agent",
|
||||
Headers: "{}",
|
||||
Status: 200,
|
||||
Latency: 12,
|
||||
CreatedAt: createdAt,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, mockConn.prepareCalled)
|
||||
assert.Equal(t, analyticsmodel.UserAccessLog{}.BatchInsertSQL(), mockConn.preparedQuery)
|
||||
assert.True(t, mockBatch.sendCalled)
|
||||
require.Len(t, mockBatch.rows, 1)
|
||||
assert.Equal(t, uint64(42), mockBatch.rows[0][1])
|
||||
}
|
||||
|
||||
type mockConn struct {
|
||||
batch driver.Batch
|
||||
batchQuery string
|
||||
prepareCalled bool
|
||||
preparedQuery string
|
||||
}
|
||||
|
||||
func (m *mockConn) Contributors() []string { return nil }
|
||||
|
||||
func (m *mockConn) ServerVersion() (*driver.ServerVersion, error) { return nil, nil }
|
||||
|
||||
func (m *mockConn) Select(_ context.Context, _ any, _ string, _ ...any) error { return nil }
|
||||
|
||||
func (m *mockConn) Query(_ context.Context, _ string, _ ...any) (driver.Rows, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *mockConn) QueryRow(_ context.Context, _ string, _ ...any) driver.Row { return nil }
|
||||
|
||||
func (m *mockConn) PrepareBatch(_ context.Context, query string, _ ...driver.PrepareBatchOption) (driver.Batch, error) {
|
||||
m.prepareCalled = true
|
||||
m.preparedQuery = query
|
||||
return m.batch, nil
|
||||
}
|
||||
|
||||
func (m *mockConn) Exec(_ context.Context, _ string, _ ...any) error { return nil }
|
||||
|
||||
func (m *mockConn) AsyncInsert(_ context.Context, _ string, _ bool, _ ...any) error { return nil }
|
||||
|
||||
func (m *mockConn) Ping(_ context.Context) error { return nil }
|
||||
|
||||
func (m *mockConn) Stats() driver.Stats { return driver.Stats{} }
|
||||
|
||||
func (m *mockConn) Close() error { return nil }
|
||||
|
||||
type mockBatch struct {
|
||||
rows [][]any
|
||||
sendCalled bool
|
||||
}
|
||||
|
||||
func (m *mockBatch) Abort() error { return nil }
|
||||
|
||||
func (m *mockBatch) Append(v ...any) error {
|
||||
m.rows = append(m.rows, v)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockBatch) AppendStruct(_ any) error { return nil }
|
||||
|
||||
func (m *mockBatch) Column(_ int) driver.BatchColumn { return nil }
|
||||
|
||||
func (m *mockBatch) Flush() error { return nil }
|
||||
|
||||
func (m *mockBatch) Send() error {
|
||||
m.sendCalled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockBatch) IsSent() bool { return m.sendCalled }
|
||||
|
||||
func (m *mockBatch) Rows() int { return len(m.rows) }
|
||||
|
||||
func (m *mockBatch) Columns() []column.Interface { return nil }
|
||||
|
||||
func (m *mockBatch) Close() error { return nil }
|
||||
@@ -0,0 +1,49 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||
)
|
||||
|
||||
// BatchInsert writes access logs to ClickHouse using the native batch API.
|
||||
func BatchInsert(ctx context.Context, logs []analyticsmodel.UserAccessLog) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
if db.ChConn == nil {
|
||||
return fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
|
||||
batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.UserAccessLog{}.BatchInsertSQL())
|
||||
if err != nil {
|
||||
return fmt.Errorf("prepare clickhouse batch: %w", err)
|
||||
}
|
||||
|
||||
for _, logItem := range logs {
|
||||
if err := batch.Append(
|
||||
logItem.ID,
|
||||
logItem.UserID,
|
||||
logItem.Path,
|
||||
logItem.Method,
|
||||
logItem.IP,
|
||||
logItem.UserAgent,
|
||||
logItem.Headers,
|
||||
logItem.Status,
|
||||
logItem.Latency,
|
||||
logItem.CreatedAt,
|
||||
); err != nil {
|
||||
return fmt.Errorf("append access log to batch: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := batch.Send(); err != nil {
|
||||
return fmt.Errorf("send clickhouse batch: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import "strings"
|
||||
|
||||
// ParseBrowserName performs lightweight User-Agent browser identification.
|
||||
func ParseBrowserName(ua string) string {
|
||||
uaLower := strings.ToLower(ua)
|
||||
if strings.Contains(uaLower, "micromessenger") {
|
||||
return "WeChat"
|
||||
}
|
||||
if strings.Contains(uaLower, "postman") {
|
||||
return "Postman"
|
||||
}
|
||||
if strings.Contains(uaLower, "edg/") || strings.Contains(uaLower, "edge") {
|
||||
return "Edge"
|
||||
}
|
||||
if strings.Contains(uaLower, "firefox") {
|
||||
return "Firefox"
|
||||
}
|
||||
if strings.Contains(uaLower, "chrome") {
|
||||
return "Chrome"
|
||||
}
|
||||
if strings.Contains(uaLower, "safari") {
|
||||
return "Safari"
|
||||
}
|
||||
return "Other"
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||
)
|
||||
|
||||
// NodeAccessLogRegionCount aggregates access log regions.
|
||||
type NodeAccessLogRegionCount struct {
|
||||
Region string
|
||||
Count int64
|
||||
}
|
||||
|
||||
func nodeAccessLogConn() (driver.Conn, error) {
|
||||
if db.ChConn == nil {
|
||||
return nil, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
return db.ChConn, nil
|
||||
}
|
||||
|
||||
// ListNodeAccessLogs returns access logs matching filter.
|
||||
func ListNodeAccessLogs(ctx context.Context, filter NodeAccessLogFilter) ([]analyticsmodel.NodeAccessLog, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
tableName := nodeAccessLogTableName()
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT id, node_id, logged_at, remote_addr, region, host, path, status_code, created_at
|
||||
FROM %s
|
||||
WHERE %s
|
||||
ORDER BY %s`, tableName, clause, nodeAccessLogOrderClause(filter.SortBy, filter.SortOrder))
|
||||
if filter.PageSize > 0 {
|
||||
if filter.Page < 0 {
|
||||
filter.Page = 0
|
||||
}
|
||||
sql += " LIMIT ? OFFSET ?"
|
||||
args = append(args, filter.PageSize, filter.Page*filter.PageSize)
|
||||
}
|
||||
rows, err := conn.Query(ctx, sql, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list node access logs: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
return scanNodeAccessLogRows(rows)
|
||||
}
|
||||
|
||||
func scanNodeAccessLogRows(rows driver.Rows) ([]analyticsmodel.NodeAccessLog, error) {
|
||||
var result []analyticsmodel.NodeAccessLog
|
||||
for rows.Next() {
|
||||
var item analyticsmodel.NodeAccessLog
|
||||
if err := rows.Scan(
|
||||
&item.ID,
|
||||
&item.NodeID,
|
||||
&item.LoggedAt,
|
||||
&item.RemoteAddr,
|
||||
&item.Region,
|
||||
&item.Host,
|
||||
&item.Path,
|
||||
&item.StatusCode,
|
||||
&item.CreatedAt,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("scan node access log row: %w", err)
|
||||
}
|
||||
item.LoggedAt = item.LoggedAt.UTC()
|
||||
item.CreatedAt = item.CreatedAt.UTC()
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// CountNodeAccessLogs returns total records and distinct IPs matching filter.
|
||||
func CountNodeAccessLogs(ctx context.Context, filter NodeAccessLogFilter) (int64, int64, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
tableName := nodeAccessLogTableName()
|
||||
|
||||
var totalRecords int64
|
||||
countSQL := fmt.Sprintf("SELECT count() FROM %s WHERE %s", tableName, clause)
|
||||
if err := conn.QueryRow(ctx, countSQL, args...).Scan(&totalRecords); err != nil {
|
||||
return 0, 0, fmt.Errorf("count node access logs: %w", err)
|
||||
}
|
||||
|
||||
ipSQL := fmt.Sprintf(`
|
||||
SELECT count() FROM (
|
||||
SELECT trim(remote_addr) AS remote_addr
|
||||
FROM %s
|
||||
WHERE %s AND remote_addr != ''
|
||||
GROUP BY trim(remote_addr)
|
||||
)`, tableName, clause)
|
||||
var totalIPs int64
|
||||
if err := conn.QueryRow(ctx, ipSQL, args...).Scan(&totalIPs); err != nil {
|
||||
return 0, 0, fmt.Errorf("count node access log ips: %w", err)
|
||||
}
|
||||
return totalRecords, totalIPs, nil
|
||||
}
|
||||
|
||||
// RegionCountsNodeAccessLogs returns region counts for a node since a time.
|
||||
func RegionCountsNodeAccessLogs(ctx context.Context, nodeID string, since time.Time, limit int) ([]NodeAccessLogRegionCount, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
filter := NodeAccessLogFilter{NodeID: nodeID, Since: since}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
tableName := nodeAccessLogTableName()
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT trim(region) AS region, count() AS count
|
||||
FROM %s
|
||||
WHERE %s AND trim(region) != ''
|
||||
GROUP BY trim(region)
|
||||
ORDER BY count DESC, region ASC`, tableName, clause)
|
||||
if limit > 0 {
|
||||
sql += " LIMIT ?"
|
||||
args = append(args, limit)
|
||||
}
|
||||
rows, err := conn.Query(ctx, sql, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("region counts node access logs: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var result []NodeAccessLogRegionCount
|
||||
for rows.Next() {
|
||||
var item NodeAccessLogRegionCount
|
||||
if err := rows.Scan(&item.Region, &item.Count); err != nil {
|
||||
return nil, fmt.Errorf("scan region count row: %w", err)
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// DeleteAllNodeAccessLogs deletes all node access logs.
|
||||
func DeleteAllNodeAccessLogs(ctx context.Context) (int64, error) {
|
||||
tableName := nodeAccessLogTableName()
|
||||
return deleteNodeAccessLogsWithCount(ctx, "SELECT count() FROM "+tableName, nil, "ALTER TABLE "+tableName+" DELETE WHERE 1")
|
||||
}
|
||||
|
||||
// DeleteNodeAccessLogsBefore deletes logs older than cutoff.
|
||||
func DeleteNodeAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
|
||||
tableName := nodeAccessLogTableName()
|
||||
cutoff = cutoff.UTC()
|
||||
return deleteNodeAccessLogsWithCount(
|
||||
ctx,
|
||||
fmt.Sprintf("SELECT count() FROM %s WHERE logged_at < ?", tableName),
|
||||
[]any{cutoff},
|
||||
fmt.Sprintf("ALTER TABLE %s DELETE WHERE logged_at < ?", tableName),
|
||||
cutoff,
|
||||
)
|
||||
}
|
||||
|
||||
// DeleteNodeAccessLogsByNodeBefore deletes logs for a node older than cutoff.
|
||||
func DeleteNodeAccessLogsByNodeBefore(ctx context.Context, nodeID string, before time.Time) (int64, error) {
|
||||
tableName := nodeAccessLogTableName()
|
||||
before = before.UTC()
|
||||
return deleteNodeAccessLogsWithCount(
|
||||
ctx,
|
||||
fmt.Sprintf("SELECT count() FROM %s WHERE node_id = ? AND logged_at < ?", tableName),
|
||||
[]any{nodeID, before},
|
||||
fmt.Sprintf("ALTER TABLE %s DELETE WHERE node_id = ? AND logged_at < ?", tableName),
|
||||
nodeID, before,
|
||||
)
|
||||
}
|
||||
|
||||
func deleteNodeAccessLogsWithCount(ctx context.Context, countSQL string, countArgs []any, deleteSQL string, deleteArgs ...any) (int64, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
var count int64
|
||||
if err := conn.QueryRow(ctx, countSQL, countArgs...).Scan(&count); err != nil {
|
||||
return 0, fmt.Errorf("count node access logs for delete: %w", err)
|
||||
}
|
||||
if count == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
if err := conn.Exec(ctx, deleteSQL, deleteArgs...); err != nil {
|
||||
return 0, fmt.Errorf("delete node access logs: %w", err)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// NodeAccessLogFilter scopes ClickHouse node access log queries.
|
||||
type NodeAccessLogFilter struct {
|
||||
NodeID string
|
||||
RemoteAddr string
|
||||
Host string
|
||||
Path string
|
||||
Since time.Time
|
||||
Until time.Time
|
||||
Page int
|
||||
PageSize int
|
||||
SortBy string
|
||||
SortOrder string
|
||||
}
|
||||
|
||||
func buildNodeAccessLogFilterClause(filter NodeAccessLogFilter) (string, []any) {
|
||||
parts := make([]string, 0, 6)
|
||||
args := make([]any, 0, 6)
|
||||
if trimmed := strings.TrimSpace(filter.NodeID); trimmed != "" {
|
||||
parts = append(parts, "node_id = ?")
|
||||
args = append(args, trimmed)
|
||||
}
|
||||
if trimmed := strings.TrimSpace(filter.RemoteAddr); trimmed != "" {
|
||||
parts = append(parts, "remote_addr LIKE ?")
|
||||
args = append(args, trimmed+"%")
|
||||
}
|
||||
if trimmed := strings.TrimSpace(filter.Host); trimmed != "" {
|
||||
parts = append(parts, "host LIKE ?")
|
||||
args = append(args, trimmed+"%")
|
||||
}
|
||||
if trimmed := strings.TrimSpace(filter.Path); trimmed != "" {
|
||||
parts = append(parts, "path LIKE ?")
|
||||
args = append(args, trimmed+"%")
|
||||
}
|
||||
if !filter.Since.IsZero() {
|
||||
parts = append(parts, "logged_at >= ?")
|
||||
args = append(args, filter.Since.UTC())
|
||||
}
|
||||
if !filter.Until.IsZero() {
|
||||
parts = append(parts, "logged_at < ?")
|
||||
args = append(args, filter.Until.UTC())
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "1", nil
|
||||
}
|
||||
return strings.Join(parts, " AND "), args
|
||||
}
|
||||
|
||||
func combineNodeAccessLogSQLClauses(left string, right string) string {
|
||||
if strings.TrimSpace(left) == "" || left == "TRUE" || left == "1" {
|
||||
return right
|
||||
}
|
||||
return left + " AND " + right
|
||||
}
|
||||
|
||||
func nodeAccessLogOrderClause(sortBy string, sortOrder string) string {
|
||||
direction := "DESC"
|
||||
if normalizeNodeAccessLogSortOrder(sortOrder) == "asc" {
|
||||
direction = "ASC"
|
||||
}
|
||||
column := "logged_at"
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "status_code":
|
||||
column = "status_code"
|
||||
case "remote_addr":
|
||||
column = "remote_addr"
|
||||
case "host":
|
||||
column = "host"
|
||||
case "path":
|
||||
column = "path"
|
||||
}
|
||||
if column == "logged_at" {
|
||||
return column + " " + direction + ", id " + direction
|
||||
}
|
||||
return column + " " + direction + ", logged_at " + direction + ", id " + direction
|
||||
}
|
||||
|
||||
func normalizeNodeAccessLogSortOrder(sortOrder string) string {
|
||||
if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") {
|
||||
return "asc"
|
||||
}
|
||||
return "desc"
|
||||
}
|
||||
|
||||
func nodeAccessLogBucketEpochExpr(bucketSeconds int64) string {
|
||||
return fmt.Sprintf("toInt64(intDiv(toUnixTimestamp(logged_at), %d) * %d)", bucketSeconds, bucketSeconds)
|
||||
}
|
||||
|
||||
func nodeAccessLogEpochExpr() string {
|
||||
return "toInt64(toUnixTimestamp(logged_at))"
|
||||
}
|
||||
|
||||
func nodeAccessLogTableName() string {
|
||||
return "of_node_access_logs"
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// NodeAccessLogBucketAggregate is a folded bucket aggregate row.
|
||||
type NodeAccessLogBucketAggregate struct {
|
||||
BucketEpoch int64
|
||||
RequestCount int64
|
||||
SuccessCount int64
|
||||
ClientErrorCount int64
|
||||
ServerErrorCount int64
|
||||
}
|
||||
|
||||
// NodeAccessLogBucketDimension is a bucket dimension value.
|
||||
type NodeAccessLogBucketDimension struct {
|
||||
BucketEpoch int64
|
||||
Value string
|
||||
}
|
||||
|
||||
// NodeAccessLogIPAggregate is an IP aggregate row.
|
||||
type NodeAccessLogIPAggregate struct {
|
||||
RemoteAddr string
|
||||
RequestCount int64
|
||||
SuccessCount int64
|
||||
ClientErrorCount int64
|
||||
ServerErrorCount int64
|
||||
LastSeenEpoch int64
|
||||
}
|
||||
|
||||
// NodeAccessLogIPSummary is an IP summary row.
|
||||
type NodeAccessLogIPSummary struct {
|
||||
RemoteAddr string
|
||||
TotalRequests int64
|
||||
RecentRequests int64
|
||||
LastSeenEpoch int64
|
||||
}
|
||||
|
||||
// NodeAccessLogIPTrend is an IP trend bucket row.
|
||||
type NodeAccessLogIPTrend struct {
|
||||
BucketEpoch int64
|
||||
RequestCount int64
|
||||
}
|
||||
|
||||
// BucketAggregatesNodeAccessLogs returns folded bucket aggregates.
|
||||
func BucketAggregatesNodeAccessLogs(ctx context.Context, filter NodeAccessLogFilter, bucketSeconds int64) ([]NodeAccessLogBucketAggregate, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
bucketExpr := nodeAccessLogBucketEpochExpr(bucketSeconds)
|
||||
tableName := nodeAccessLogTableName()
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT
|
||||
%s AS bucket_epoch,
|
||||
count() AS request_count,
|
||||
countIf(status_code < 400) AS success_count,
|
||||
countIf(status_code >= 400 AND status_code < 500) AS client_error_count,
|
||||
countIf(status_code >= 500) AS server_error_count
|
||||
FROM %s
|
||||
WHERE %s
|
||||
GROUP BY bucket_epoch`, bucketExpr, tableName, clause)
|
||||
rows, err := conn.Query(ctx, sql, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bucket aggregates node access logs: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var result []NodeAccessLogBucketAggregate
|
||||
for rows.Next() {
|
||||
var item NodeAccessLogBucketAggregate
|
||||
if err := rows.Scan(&item.BucketEpoch, &item.RequestCount, &item.SuccessCount, &item.ClientErrorCount, &item.ServerErrorCount); err != nil {
|
||||
return nil, fmt.Errorf("scan bucket aggregate row: %w", err)
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// BucketDimensionsNodeAccessLogs returns bucket dimension values.
|
||||
func BucketDimensionsNodeAccessLogs(ctx context.Context, filter NodeAccessLogFilter, column string, bucketSeconds int64) ([]NodeAccessLogBucketDimension, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
bucketExpr := nodeAccessLogBucketEpochExpr(bucketSeconds)
|
||||
tableName := nodeAccessLogTableName()
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT
|
||||
%s AS bucket_epoch,
|
||||
trim(%s) AS value
|
||||
FROM %s
|
||||
WHERE %s AND trim(%s) != ''
|
||||
GROUP BY bucket_epoch, trim(%s)`, bucketExpr, column, tableName, clause, column, column)
|
||||
rows, err := conn.Query(ctx, sql, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bucket dimensions node access logs: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var result []NodeAccessLogBucketDimension
|
||||
for rows.Next() {
|
||||
var item NodeAccessLogBucketDimension
|
||||
if err := rows.Scan(&item.BucketEpoch, &item.Value); err != nil {
|
||||
return nil, fmt.Errorf("scan bucket dimension row: %w", err)
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// IPAggregatesNodeAccessLogs returns IP aggregate rows.
|
||||
func IPAggregatesNodeAccessLogs(ctx context.Context, filter NodeAccessLogFilter, exactRemoteAddr bool) ([]NodeAccessLogIPAggregate, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
queryClause := clause
|
||||
queryArgs := append([]any{}, args...)
|
||||
if exactRemoteAddr {
|
||||
trimmed := strings.TrimSpace(filter.RemoteAddr)
|
||||
if trimmed == "" {
|
||||
return []NodeAccessLogIPAggregate{}, nil
|
||||
}
|
||||
queryClause = combineNodeAccessLogSQLClauses(queryClause, "trim(remote_addr) = ?")
|
||||
queryArgs = append(queryArgs, trimmed)
|
||||
}
|
||||
lastSeenExpr := nodeAccessLogEpochExpr()
|
||||
tableName := nodeAccessLogTableName()
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT
|
||||
trim(remote_addr) AS remote_addr,
|
||||
count() AS request_count,
|
||||
countIf(status_code < 400) AS success_count,
|
||||
countIf(status_code >= 400 AND status_code < 500) AS client_error_count,
|
||||
countIf(status_code >= 500) AS server_error_count,
|
||||
max(%s) AS last_seen_epoch
|
||||
FROM %s
|
||||
WHERE %s AND trim(remote_addr) != ''
|
||||
GROUP BY trim(remote_addr)`, lastSeenExpr, tableName, queryClause)
|
||||
rows, err := conn.Query(ctx, sql, queryArgs...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ip aggregates node access logs: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var result []NodeAccessLogIPAggregate
|
||||
for rows.Next() {
|
||||
var item NodeAccessLogIPAggregate
|
||||
if err := rows.Scan(&item.RemoteAddr, &item.RequestCount, &item.SuccessCount, &item.ClientErrorCount, &item.ServerErrorCount, &item.LastSeenEpoch); err != nil {
|
||||
return nil, fmt.Errorf("scan ip aggregate row: %w", err)
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// IPSummariesNodeAccessLogs returns IP summary rows.
|
||||
func IPSummariesNodeAccessLogs(ctx context.Context, filter NodeAccessLogFilter, recentSince time.Time) ([]NodeAccessLogIPSummary, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
lastSeenExpr := nodeAccessLogEpochExpr()
|
||||
recentClause := "0"
|
||||
queryArgs := make([]any, 0, len(args)+1)
|
||||
if !recentSince.IsZero() {
|
||||
recentClause = "if(logged_at >= ?, 1, 0)"
|
||||
queryArgs = append(queryArgs, recentSince)
|
||||
}
|
||||
queryArgs = append(queryArgs, args...)
|
||||
tableName := nodeAccessLogTableName()
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT
|
||||
trim(remote_addr) AS remote_addr,
|
||||
count() AS total_requests,
|
||||
sum(%s) AS recent_requests,
|
||||
max(%s) AS last_seen_epoch
|
||||
FROM %s
|
||||
WHERE %s AND trim(remote_addr) != ''
|
||||
GROUP BY trim(remote_addr)`, recentClause, lastSeenExpr, tableName, clause)
|
||||
rows, err := conn.Query(ctx, sql, queryArgs...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ip summaries node access logs: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var result []NodeAccessLogIPSummary
|
||||
for rows.Next() {
|
||||
var item NodeAccessLogIPSummary
|
||||
if err := rows.Scan(&item.RemoteAddr, &item.TotalRequests, &item.RecentRequests, &item.LastSeenEpoch); err != nil {
|
||||
return nil, fmt.Errorf("scan ip summary row: %w", err)
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// IPTrendNodeAccessLogs returns IP trend bucket rows.
|
||||
func IPTrendNodeAccessLogs(ctx context.Context, filter NodeAccessLogFilter, bucketSeconds int64) ([]NodeAccessLogIPTrend, error) {
|
||||
conn, err := nodeAccessLogConn()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clause, args := buildNodeAccessLogFilterClause(filter)
|
||||
bucketExpr := nodeAccessLogBucketEpochExpr(bucketSeconds)
|
||||
tableName := nodeAccessLogTableName()
|
||||
sql := fmt.Sprintf(`
|
||||
SELECT
|
||||
%s AS bucket_epoch,
|
||||
count() AS request_count
|
||||
FROM %s
|
||||
WHERE %s
|
||||
GROUP BY bucket_epoch
|
||||
ORDER BY bucket_epoch ASC`, bucketExpr, tableName, clause)
|
||||
rows, err := conn.Query(ctx, sql, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ip trend node access logs: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var result []NodeAccessLogIPTrend
|
||||
for rows.Next() {
|
||||
var item NodeAccessLogIPTrend
|
||||
if err := rows.Scan(&item.BucketEpoch, &item.RequestCount); err != nil {
|
||||
return nil, fmt.Errorf("scan ip trend row: %w", err)
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestBatchInsertNodeAccessLogs_Empty(t *testing.T) {
|
||||
err := BatchInsertNodeAccessLogs(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestBatchInsertNodeAccessLogs_UsesModelBatchSQL(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mockBatch := &mockBatch{}
|
||||
mockConn := &mockConn{
|
||||
batch: mockBatch,
|
||||
batchQuery: analyticsmodel.NodeAccessLog{}.BatchInsertSQL(),
|
||||
}
|
||||
db.SetChConnForTest(mockConn)
|
||||
t.Cleanup(func() { db.SetChConnForTest(nil) })
|
||||
|
||||
loggedAt := time.Now().UTC()
|
||||
err := BatchInsertNodeAccessLogs(ctx, []analyticsmodel.NodeAccessLog{
|
||||
{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: loggedAt,
|
||||
RemoteAddr: "1.1.1.1",
|
||||
Region: "US",
|
||||
Host: "example.com",
|
||||
Path: "/alpha",
|
||||
StatusCode: 200,
|
||||
CreatedAt: loggedAt,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, mockConn.prepareCalled)
|
||||
assert.Equal(t, analyticsmodel.NodeAccessLog{}.BatchInsertSQL(), mockConn.preparedQuery)
|
||||
assert.True(t, mockBatch.sendCalled)
|
||||
require.Len(t, mockBatch.rows, 1)
|
||||
assert.Equal(t, "node-a", mockBatch.rows[0][1])
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||
)
|
||||
|
||||
// BatchInsertNodeAccessLogs writes node access logs to ClickHouse using the native batch API.
|
||||
func BatchInsertNodeAccessLogs(ctx context.Context, logs []analyticsmodel.NodeAccessLog) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
if db.ChConn == nil {
|
||||
return fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
|
||||
batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeAccessLog{}.BatchInsertSQL())
|
||||
if err != nil {
|
||||
return fmt.Errorf("prepare clickhouse batch: %w", err)
|
||||
}
|
||||
|
||||
now := time.Now().UTC()
|
||||
for _, logItem := range logs {
|
||||
id := logItem.ID
|
||||
if id == 0 {
|
||||
id = idgen.NextUint64ID()
|
||||
}
|
||||
createdAt := logItem.CreatedAt
|
||||
if createdAt.IsZero() {
|
||||
createdAt = now
|
||||
}
|
||||
if err := batch.Append(
|
||||
id,
|
||||
logItem.NodeID,
|
||||
logItem.LoggedAt.UTC(),
|
||||
logItem.RemoteAddr,
|
||||
logItem.Region,
|
||||
logItem.Host,
|
||||
logItem.Path,
|
||||
logItem.StatusCode,
|
||||
createdAt.UTC(),
|
||||
); err != nil {
|
||||
return fmt.Errorf("append node access log to batch: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := batch.Send(); err != nil {
|
||||
return fmt.Errorf("send clickhouse batch: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const activePushChannelCacheTTL = 24 * time.Hour
|
||||
|
||||
// ListPushChannels returns all push channels ordered by creation time descending.
|
||||
func ListPushChannels(ctx context.Context) ([]model.PushChannel, error) {
|
||||
var channels []model.PushChannel
|
||||
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return channels, nil
|
||||
}
|
||||
|
||||
// GetPushChannelByID loads a push channel by primary key.
|
||||
func GetPushChannelByID(ctx context.Context, id uint64) (model.PushChannel, error) {
|
||||
var channel model.PushChannel
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
||||
return model.PushChannel{}, err
|
||||
}
|
||||
return channel, nil
|
||||
}
|
||||
|
||||
// GetPushChannelByName 根据名称获取消息通道。
|
||||
func GetPushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) {
|
||||
var channel model.PushChannel
|
||||
if err := db.DB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &channel, nil
|
||||
}
|
||||
|
||||
// CountPushChannelsByName returns how many channels share the given name.
|
||||
func CountPushChannelsByName(ctx context.Context, name string) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CreatePushChannel persists a new channel and invalidates cache.
|
||||
func CreatePushChannel(ctx context.Context, channel *model.PushChannel) error {
|
||||
if err := db.DB(ctx).Create(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SavePushChannel updates a channel and invalidates cache.
|
||||
func SavePushChannel(ctx context.Context, channel *model.PushChannel) error {
|
||||
if err := db.DB(ctx).Save(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeletePushChannel removes a channel and invalidates cache.
|
||||
func DeletePushChannel(ctx context.Context, channel *model.PushChannel) error {
|
||||
if err := db.DB(ctx).Delete(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
|
||||
func GetActivePushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) {
|
||||
cacheKey := "push:channel:active:" + name
|
||||
var channel model.PushChannel
|
||||
if db.Redis != nil {
|
||||
if err := db.GetJSON(ctx, cacheKey, &channel); err == nil {
|
||||
return &channel, nil
|
||||
}
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if db.Redis != nil {
|
||||
_ = db.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||
}
|
||||
|
||||
return &channel, nil
|
||||
}
|
||||
|
||||
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
|
||||
func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey("push:channel:active:"+name)).Err()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const activePushEventCacheTTL = 24 * time.Hour
|
||||
|
||||
// ListPushEvents returns all push events ordered by creation time descending.
|
||||
func ListPushEvents(ctx context.Context) ([]model.PushEvent, error) {
|
||||
var events []model.PushEvent
|
||||
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return events, nil
|
||||
}
|
||||
|
||||
// GetPushEventByID loads a push event by primary key.
|
||||
func GetPushEventByID(ctx context.Context, id uint64) (model.PushEvent, error) {
|
||||
var event model.PushEvent
|
||||
if err := db.DB(ctx).First(&event, id).Error; err != nil {
|
||||
return model.PushEvent{}, err
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
|
||||
// GetPushEventByKey loads a push event by event key.
|
||||
func GetPushEventByKey(ctx context.Context, key string) (model.PushEvent, error) {
|
||||
var event model.PushEvent
|
||||
if err := db.DB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
|
||||
return model.PushEvent{}, err
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
|
||||
// CountPushEventsByKey returns how many events use the given event key.
|
||||
func CountPushEventsByKey(ctx context.Context, key string) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CreatePushEvent persists a new push event and invalidates cache.
|
||||
func CreatePushEvent(ctx context.Context, event *model.PushEvent) error {
|
||||
if err := db.DB(ctx).Create(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SavePushEvent updates a push event and invalidates cache.
|
||||
func SavePushEvent(ctx context.Context, event *model.PushEvent) error {
|
||||
if err := db.DB(ctx).Save(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdatePushEventEnabled toggles the enabled flag for a push event.
|
||||
func UpdatePushEventEnabled(ctx context.Context, event *model.PushEvent, enabled bool) error {
|
||||
event.Enabled = enabled
|
||||
if err := db.DB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeletePushEvent removes a push event and invalidates cache.
|
||||
func DeletePushEvent(ctx context.Context, event *model.PushEvent) error {
|
||||
if err := db.DB(ctx).Delete(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListActivePushEventsByTaskType returns enabled events bound to a task type.
|
||||
func ListActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) {
|
||||
var events []model.PushEvent
|
||||
if err := db.DB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return events, nil
|
||||
}
|
||||
|
||||
// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)。
|
||||
func GetActivePushEventByKey(ctx context.Context, key string) (*model.PushEvent, error) {
|
||||
cacheKey := "push:event:active:" + key
|
||||
var event model.PushEvent
|
||||
if db.Redis != nil {
|
||||
if err := db.GetJSON(ctx, cacheKey, &event); err == nil {
|
||||
return &event, nil
|
||||
}
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if db.Redis != nil {
|
||||
_ = db.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||
}
|
||||
|
||||
return &event, nil
|
||||
}
|
||||
|
||||
// DeleteActivePushEventCache 清理启用通知事件的缓存。
|
||||
func DeleteActivePushEventCache(ctx context.Context, key string) {
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey("push:event:active:"+key)).Err()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// PushHistoryListFilter filters push history pagination queries.
|
||||
type PushHistoryListFilter struct {
|
||||
EventKey string
|
||||
Status string
|
||||
Page int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
// ListPushHistories returns paginated push history records.
|
||||
func ListPushHistories(ctx context.Context, filter PushHistoryListFilter) (int64, []model.PushHistory, error) {
|
||||
query := db.DB(ctx).Model(&model.PushHistory{}).Order("created_at DESC")
|
||||
if filter.EventKey != "" {
|
||||
query = query.Where("event_key = ?", filter.EventKey)
|
||||
}
|
||||
if filter.Status != "" {
|
||||
query = query.Where("status = ?", filter.Status)
|
||||
}
|
||||
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
var results []model.PushHistory
|
||||
offset := (filter.Page - 1) * filter.PageSize
|
||||
if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
return total, results, nil
|
||||
}
|
||||
|
||||
// CreatePushHistory persists a push history audit record.
|
||||
func CreatePushHistory(ctx context.Context, history *model.PushHistory) error {
|
||||
return db.DB(ctx).Create(history).Error
|
||||
}
|
||||
|
||||
// PushHistoryQuery returns a scoped query builder for push histories.
|
||||
func PushHistoryQuery(ctx context.Context) *gorm.DB {
|
||||
return db.DB(ctx).Model(&model.PushHistory{})
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package repository provides data access with caching and persistence boundaries.
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/shopspring/decimal"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
errDatabaseNotInitialized = "database not initialized"
|
||||
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
|
||||
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
|
||||
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
|
||||
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
|
||||
)
|
||||
|
||||
// GetSystemConfigByKey 通过 key 查询配置(带 RAM + Redis 缓存)。
|
||||
func GetSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
|
||||
ensureSystemConfigCacheListener()
|
||||
|
||||
if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok {
|
||||
return cloneSystemConfig(cached), nil
|
||||
}
|
||||
|
||||
var sc model.SystemConfig
|
||||
if db.Redis != nil {
|
||||
if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, &sc); err == nil {
|
||||
systemConfigRAMCache.Set(key, cloneSystemConfig(sc))
|
||||
return sc, nil
|
||||
} else if !errors.Is(err, redis.Nil) {
|
||||
return model.SystemConfig{}, err
|
||||
}
|
||||
}
|
||||
|
||||
database := db.DB(ctx)
|
||||
if database == nil {
|
||||
return model.SystemConfig{}, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
|
||||
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
|
||||
return model.SystemConfig{}, err
|
||||
}
|
||||
|
||||
populateSystemConfigCache(ctx, sc)
|
||||
return sc, nil
|
||||
}
|
||||
|
||||
// ListSystemConfigsByKeys loads multiple config keys in one database round trip.
|
||||
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]model.SystemConfig, error) {
|
||||
if len(keys) == 0 {
|
||||
return map[string]model.SystemConfig{}, nil
|
||||
}
|
||||
|
||||
ensureSystemConfigCacheListener()
|
||||
|
||||
result := make(map[string]model.SystemConfig, len(keys))
|
||||
missing := make([]string, 0, len(keys))
|
||||
for _, key := range keys {
|
||||
if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok {
|
||||
result[key] = cloneSystemConfig(cached)
|
||||
continue
|
||||
}
|
||||
missing = append(missing, key)
|
||||
}
|
||||
|
||||
if len(missing) == 0 {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
database := db.DB(ctx)
|
||||
if database == nil {
|
||||
return nil, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
|
||||
var configs []model.SystemConfig
|
||||
if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for i := range configs {
|
||||
populateSystemConfigCache(ctx, configs[i])
|
||||
result[configs[i].Key] = cloneSystemConfig(configs[i])
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
|
||||
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
|
||||
if db.Redis == nil {
|
||||
return nil
|
||||
}
|
||||
return db.Redis.Del(ctx, db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
|
||||
}
|
||||
|
||||
// ListVisibleSystemConfigs 查询所有可通过公共配置接口暴露的配置(带 Redis 列表缓存)。
|
||||
func ListVisibleSystemConfigs(ctx context.Context) ([]model.SystemConfig, error) {
|
||||
if db.Redis != nil {
|
||||
var cached []model.SystemConfig
|
||||
if err := db.GetJSON(ctx, SystemConfigVisibleListRedisKey, &cached); err == nil {
|
||||
return cached, nil
|
||||
} else if !errors.Is(err, redis.Nil) {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
database := db.DB(ctx)
|
||||
if database == nil {
|
||||
return nil, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
|
||||
var configs []model.SystemConfig
|
||||
if err := database.Where("visibility = ?", model.ConfigVisibilityVisible).Find(&configs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if db.Redis != nil {
|
||||
_ = db.SetJSON(ctx, SystemConfigVisibleListRedisKey, configs, 0)
|
||||
}
|
||||
|
||||
return configs, nil
|
||||
}
|
||||
|
||||
// GetIntByKey 通过 key 查询配置并转换为 int 类型。
|
||||
func GetIntByKey(ctx context.Context, key string) (int, error) {
|
||||
sc, err := GetSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
value, err := strconv.Atoi(sc.Value)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
|
||||
}
|
||||
|
||||
return value, nil
|
||||
}
|
||||
|
||||
// GetDecimalByKey 通过 key 查询配置并转换为 decimal.Decimal 类型。
|
||||
func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) {
|
||||
sc, err := GetSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return decimal.Zero, err
|
||||
}
|
||||
|
||||
value, err := decimal.NewFromString(sc.Value)
|
||||
if err != nil {
|
||||
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
|
||||
}
|
||||
|
||||
return value.Truncate(precision), nil
|
||||
}
|
||||
|
||||
// GetBoolByKey 通过 key 查询配置并转换为 bool 类型。
|
||||
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
|
||||
sc, err := GetSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
value, err := strconv.ParseBool(sc.Value)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
|
||||
}
|
||||
|
||||
return value, nil
|
||||
}
|
||||
|
||||
// GetMenuDisplayConfig 获取目录显示配置,解析为 map[string]bool。
|
||||
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
|
||||
sc, err := GetSystemConfigByKey(ctx, model.ConfigKeyMenuDisplayConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
config := make(map[string]bool)
|
||||
if sc.Value == "" || sc.Value == "{}" {
|
||||
return config, nil
|
||||
}
|
||||
|
||||
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
|
||||
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
|
||||
}
|
||||
|
||||
return config, nil
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
|
||||
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
|
||||
query := db.DB(ctx).Order("created_at DESC")
|
||||
if configType != "" {
|
||||
query = query.Where("type = ?", configType)
|
||||
}
|
||||
var configs []model.SystemConfig
|
||||
if err := query.Find(&configs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return configs, nil
|
||||
}
|
||||
|
||||
// GetAdminSystemConfigByKey loads a config directly from PostgreSQL.
|
||||
func GetAdminSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
|
||||
var config model.SystemConfig
|
||||
if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
|
||||
return model.SystemConfig{}, err
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// SystemConfigExists reports whether a config key already exists.
|
||||
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
|
||||
var existing model.SystemConfig
|
||||
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// CreateSystemConfig persists a new system config row.
|
||||
func CreateSystemConfig(ctx context.Context, config *model.SystemConfig) error {
|
||||
return db.DB(ctx).Create(config).Error
|
||||
}
|
||||
|
||||
// UpdateSystemConfigFields applies partial updates to a system config row.
|
||||
func UpdateSystemConfigFields(ctx context.Context, config *model.SystemConfig, updates map[string]any) error {
|
||||
return db.DB(ctx).Model(config).Updates(updates).Error
|
||||
}
|
||||
|
||||
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
|
||||
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
|
||||
var sc model.SystemConfig
|
||||
err := db.DB(ctx).Where("key = ?", key).First(&sc).Error
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return err
|
||||
}
|
||||
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
sc = model.SystemConfig{
|
||||
Key: key,
|
||||
Value: value,
|
||||
Type: "system",
|
||||
Visibility: model.ConfigVisibilityHidden,
|
||||
}
|
||||
if err := db.DB(ctx).Create(&sc).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
sc.Value = value
|
||||
if err := db.DB(ctx).Save(&sc).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return InvalidateSystemConfigCache(ctx, key)
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"sync"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||
)
|
||||
|
||||
const (
|
||||
// SystemConfigInvalidationChannel broadcasts RAM cache eviction across nodes.
|
||||
SystemConfigInvalidationChannel = "system:config_invalidation"
|
||||
// SystemConfigRedisHashKey Redis Hash key,存储所有系统配置。
|
||||
SystemConfigRedisHashKey = "system:system_configs"
|
||||
// SystemConfigVisibleListRedisKey Redis key,缓存所有 visibility=1 的公共配置列表。
|
||||
SystemConfigVisibleListRedisKey = "system:visible_configs"
|
||||
|
||||
systemConfigInvalidateAllToken = "*"
|
||||
systemConfigRAMMaximumSize = 512
|
||||
)
|
||||
|
||||
type systemConfigInvalidationMessage struct {
|
||||
Key string `json:"key"`
|
||||
}
|
||||
|
||||
var (
|
||||
systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize})
|
||||
systemConfigListenerOnce sync.Once
|
||||
)
|
||||
|
||||
func ensureSystemConfigCacheListener() {
|
||||
systemConfigListenerOnce.Do(startSystemConfigCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startSystemConfigCacheInvalidationListener() {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
pubsub := db.Redis.Subscribe(context.Background(), SystemConfigInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
var payload systemConfigInvalidationMessage
|
||||
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
|
||||
systemConfigRAMCache.InvalidateAll()
|
||||
continue
|
||||
}
|
||||
if payload.Key == "" || payload.Key == systemConfigInvalidateAllToken {
|
||||
systemConfigRAMCache.InvalidateAll()
|
||||
continue
|
||||
}
|
||||
systemConfigRAMCache.Invalidate(payload.Key)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func cloneSystemConfig(sc model.SystemConfig) model.SystemConfig {
|
||||
return sc
|
||||
}
|
||||
|
||||
func populateSystemConfigCache(ctx context.Context, sc model.SystemConfig) {
|
||||
systemConfigRAMCache.Set(sc.Key, cloneSystemConfig(sc))
|
||||
if db.Redis != nil {
|
||||
_ = db.HSetJSON(ctx, SystemConfigRedisHashKey, sc.Key, &sc)
|
||||
}
|
||||
}
|
||||
|
||||
func publishSystemConfigRAMInvalidation(ctx context.Context, key string) {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
payload, err := json.Marshal(systemConfigInvalidationMessage{Key: key})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Publish(ctx, SystemConfigInvalidationChannel, payload).Err()
|
||||
}
|
||||
|
||||
// InvalidateSystemConfigCache evicts one config key from local RAM and Redis.
|
||||
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
||||
ensureSystemConfigCacheListener()
|
||||
|
||||
systemConfigRAMCache.Invalidate(key)
|
||||
if db.Redis != nil {
|
||||
if err := db.HDel(ctx, SystemConfigRedisHashKey, key); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
publishSystemConfigRAMInvalidation(ctx, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// InvalidateAllSystemConfigCaches evicts all config entries from local RAM and Redis.
|
||||
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
|
||||
ensureSystemConfigCacheListener()
|
||||
|
||||
systemConfigRAMCache.InvalidateAll()
|
||||
if db.Redis != nil {
|
||||
if err := db.Redis.Del(ctx, db.PrefixedKey(SystemConfigRedisHashKey)).Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
publishSystemConfigRAMInvalidation(ctx, systemConfigInvalidateAllToken)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
|
||||
func ResetSystemConfigRAMCacheForTest() {
|
||||
systemConfigRAMCache.InvalidateAll()
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// ListTemplates returns all templates ordered by system flag and creation time.
|
||||
func ListTemplates(ctx context.Context) ([]model.Template, error) {
|
||||
var templates []model.Template
|
||||
if err := db.DB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return templates, nil
|
||||
}
|
||||
|
||||
// GetTemplateByKey loads a template by its key.
|
||||
func GetTemplateByKey(ctx context.Context, key string) (model.Template, error) {
|
||||
var tmpl model.Template
|
||||
if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
// TemplateExistsByKey reports whether a template key is already taken.
|
||||
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
|
||||
var existing model.Template
|
||||
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// CreateTemplate persists a new template.
|
||||
func CreateTemplate(ctx context.Context, tmpl *model.Template) error {
|
||||
return db.DB(ctx).Create(tmpl).Error
|
||||
}
|
||||
|
||||
// SaveTemplate updates an existing template.
|
||||
func SaveTemplate(ctx context.Context, tmpl *model.Template) error {
|
||||
return db.DB(ctx).Save(tmpl).Error
|
||||
}
|
||||
|
||||
// DeleteTemplate removes a template record.
|
||||
func DeleteTemplate(ctx context.Context, tmpl *model.Template) error {
|
||||
return db.DB(ctx).Delete(tmpl).Error
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// UploadListFilter filters paginated upload queries.
|
||||
type UploadListFilter struct {
|
||||
UserID uint64
|
||||
Keyword string
|
||||
Type string
|
||||
Extension string
|
||||
Page int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
// ListUploads returns paginated upload records matching the filter.
|
||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []model.Upload, error) {
|
||||
query := db.DB(ctx).Model(&model.Upload{}).
|
||||
Where("status != ?", model.UploadStatusDeleted)
|
||||
|
||||
if filter.UserID != 0 {
|
||||
query = query.Where("user_id = ?", filter.UserID)
|
||||
}
|
||||
if filter.Keyword != "" {
|
||||
query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(filter.Keyword)+"%")
|
||||
}
|
||||
if filter.Type != "" {
|
||||
query = query.Where("type = ?", filter.Type)
|
||||
}
|
||||
if filter.Extension != "" {
|
||||
query = query.Where("extension = ?", strings.ToLower(filter.Extension))
|
||||
}
|
||||
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
var items []model.Upload
|
||||
offset := (filter.Page - 1) * filter.PageSize
|
||||
if err := query.Order("created_at DESC").Offset(offset).Limit(filter.PageSize).Find(&items).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
return total, items, nil
|
||||
}
|
||||
|
||||
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||
func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
|
||||
var upload model.Upload
|
||||
if err := db.DB(ctx).Where("id = ? AND status != ?", id, model.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
return model.Upload{}, err
|
||||
}
|
||||
return upload, nil
|
||||
}
|
||||
|
||||
// SoftDeleteUpload marks an upload as deleted.
|
||||
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
|
||||
func SoftDeleteUpload(ctx context.Context, upload *model.Upload) error {
|
||||
return db.DB(ctx).Model(upload).Update("status", model.UploadStatusDeleted).Error
|
||||
}
|
||||
|
||||
// UpdateUpload applies partial field updates to an upload record.
|
||||
func UpdateUpload(ctx context.Context, upload *model.Upload, updates map[string]any) error {
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
return db.DB(ctx).Model(upload).Updates(updates).Error
|
||||
}
|
||||
|
||||
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
var types []string
|
||||
if err := db.DB(ctx).Model(&model.Upload{}).
|
||||
Where("type IS NOT NULL AND type != ''").
|
||||
Distinct().
|
||||
Pluck("type", &types).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return types, nil
|
||||
}
|
||||
|
||||
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (model.Upload, error) {
|
||||
var existing model.Upload
|
||||
err := db.DB(ctx).
|
||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, model.UploadStatusPending, model.UploadStatusUsed).
|
||||
First(&existing).Error
|
||||
return existing, err
|
||||
}
|
||||
|
||||
// CreateUpload persists a new upload record.
|
||||
// External modules must use upload.Ingest; only internal/apps/upload may call this.
|
||||
func CreateUpload(ctx context.Context, upload *model.Upload) error {
|
||||
return db.DB(ctx).Create(upload).Error
|
||||
}
|
||||
|
||||
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]model.Upload, error) {
|
||||
var uploads []model.Upload
|
||||
if err := db.DB(ctx).
|
||||
Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed).
|
||||
Find(&uploads).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return uploads, nil
|
||||
}
|
||||
|
||||
// UploadQuery returns a scoped GORM query for uploads.
|
||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||
return db.DB(ctx).Model(&model.Upload{})
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// ListUploadStats returns all upload statistics rows.
|
||||
func ListUploadStats(ctx context.Context) ([]model.UploadStat, error) {
|
||||
var stats []model.UploadStat
|
||||
if err := db.DB(ctx).Find(&stats).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return stats, nil
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// GetUserByID loads an active user by ID.
|
||||
func GetUserByID(ctx context.Context, id uint64) (model.User, error) {
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// GetUserByUsername loads a user by username.
|
||||
func GetUserByUsername(ctx context.Context, username string) (model.User, error) {
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("username = ?", username).First(&user).Error; err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// GetSystemUser loads the built-in system user, or returns a synthetic fallback.
|
||||
func GetSystemUser(ctx context.Context) model.User {
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("username = ?", "system").First(&user).Error; err == nil {
|
||||
return user
|
||||
}
|
||||
return model.User{
|
||||
ID: 999,
|
||||
Username: "system",
|
||||
Nickname: "系统",
|
||||
}
|
||||
}
|
||||
|
||||
// GetFirstAdminUser loads the earliest admin user.
|
||||
func GetFirstAdminUser(ctx context.Context) (model.User, error) {
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("is_admin = ?", true).Order("id asc").First(&user).Error; err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// AdminUserListFilter filters admin user list queries.
|
||||
type AdminUserListFilter struct {
|
||||
UserID *uint64
|
||||
Username string
|
||||
Page int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
// ListAdminUsers returns paginated users for the admin console.
|
||||
func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []model.User, error) {
|
||||
query := db.DB(ctx).Model(&model.User{})
|
||||
if filter.UserID != nil {
|
||||
query = query.Where("id = ?", *filter.UserID)
|
||||
}
|
||||
if filter.Username != "" {
|
||||
query = query.Where("username LIKE ?", filter.Username+"%")
|
||||
}
|
||||
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
var users []model.User
|
||||
offset := (filter.Page - 1) * filter.PageSize
|
||||
if err := query.
|
||||
Select("id, username, nickname, avatar_url, is_active, is_admin, last_login_at, created_at, updated_at").
|
||||
Order("id ASC").
|
||||
Offset(offset).
|
||||
Limit(filter.PageSize).
|
||||
Find(&users).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
return total, users, nil
|
||||
}
|
||||
|
||||
// GetAdminUserDetail loads full user profile fields for admin detail view.
|
||||
func GetAdminUserDetail(ctx context.Context, id uint64) (model.User, error) {
|
||||
var user model.User
|
||||
if err := db.DB(ctx).
|
||||
Select("id, username, nickname, email, avatar_url, is_active, is_admin, bio, phone, gender, website, location, last_login_at, created_at, updated_at").
|
||||
Where("id = ?", id).
|
||||
First(&user).Error; err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// UserAdminFlags stores minimal user authorization flags.
|
||||
type UserAdminFlags struct {
|
||||
ID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
|
||||
// GetUserAdminFlags loads id and is_admin for authorization checks.
|
||||
func GetUserAdminFlags(ctx context.Context, id uint64) (UserAdminFlags, error) {
|
||||
var flags UserAdminFlags
|
||||
if err := db.DB(ctx).
|
||||
Model(&model.User{}).
|
||||
Select("id, is_admin").
|
||||
Where("id = ?", id).
|
||||
First(&flags).Error; err != nil {
|
||||
return UserAdminFlags{}, err
|
||||
}
|
||||
return flags, nil
|
||||
}
|
||||
|
||||
// UpdateUserActive updates the is_active flag for a user.
|
||||
func UpdateUserActive(ctx context.Context, id uint64, active bool) error {
|
||||
return db.DB(ctx).Model(&model.User{}).Where("id = ?", id).Update("is_active", active).Error
|
||||
}
|
||||
|
||||
// DeleteUserWithRelations removes a user and related access tokens / external accounts.
|
||||
func DeleteUserWithRelations(ctx context.Context, id uint64) error {
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("user_id = ?", id).Delete(&model.AccessToken{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("user_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Where("id = ?", id).Delete(&model.User{}).Error
|
||||
})
|
||||
}
|
||||
|
||||
// CountUsersByUsername returns how many users share the username.
|
||||
func CountUsersByUsername(ctx context.Context, username string) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", username).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CountUsersByEmail returns how many users share the email.
|
||||
func CountUsersByEmail(ctx context.Context, email string) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CreateUser persists a new user record.
|
||||
func CreateUser(ctx context.Context, user *model.User) error {
|
||||
return db.DB(ctx).Create(user).Error
|
||||
}
|
||||
|
||||
// ListUsersByIDs loads users matching the given IDs.
|
||||
func ListUsersByIDs(ctx context.Context, ids []uint64) ([]model.User, error) {
|
||||
if len(ids) == 0 {
|
||||
return []model.User{}, nil
|
||||
}
|
||||
var users []model.User
|
||||
if err := db.DB(ctx).Where("id IN ?", ids).Find(&users).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return users, nil
|
||||
}
|
||||
Reference in New Issue
Block a user