Merge pull request #1512 from trheyi/main

feat: add PostgreSQL support; implement SUI for Agent with Markdown like /some/route.md
This commit is contained in:
Max 2026-04-06 18:14:40 +08:00 committed by GitHub
commit 5a69d6f595
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
29 changed files with 1436 additions and 620 deletions

View file

@ -1227,7 +1227,7 @@ jobs:
strategy: strategy:
matrix: matrix:
go: ["1.25"] go: ["1.25"]
db: [MySQL8.0, SQLite3] db: [MySQL8.0, SQLite3, Postgres14.0]
if: > if: >
${{ github.event.workflow_run.event == 'pull_request' && ${{ github.event.workflow_run.event == 'pull_request' &&
github.event.workflow_run.conclusion == 'success' }} github.event.workflow_run.conclusion == 'success' }}
@ -1358,6 +1358,8 @@ jobs:
echo "YAO_DB_DRIVER=$DB_DRIVER" >> $GITHUB_ENV echo "YAO_DB_DRIVER=$DB_DRIVER" >> $GITHUB_ENV
if [ "$DB_DRIVER" = "mysql" ]; then if [ "$DB_DRIVER" = "mysql" ]; then
echo "YAO_DB_PRIMARY=$DB_USER:$PASSWORD@$DB_HOST" >> $GITHUB_ENV echo "YAO_DB_PRIMARY=$DB_USER:$PASSWORD@$DB_HOST" >> $GITHUB_ENV
elif [ "$DB_DRIVER" = "postgres" ]; then
echo "YAO_DB_PRIMARY=postgres://$DB_USER:$PASSWORD@$DB_HOST" >> $GITHUB_ENV
else else
echo "YAO_DB_PRIMARY=${{ github.WORKSPACE }}/../app/db/yao.db" >> $GITHUB_ENV echo "YAO_DB_PRIMARY=${{ github.WORKSPACE }}/../app/db/yao.db" >> $GITHUB_ENV
mkdir -p ${{ github.WORKSPACE }}/../app/db mkdir -p ${{ github.WORKSPACE }}/../app/db
@ -1404,8 +1406,8 @@ jobs:
strategy: strategy:
matrix: matrix:
go: ["1.25"] go: ["1.25"]
db: [MySQL8.0, SQLite3] db: [MySQL8.0, SQLite3, Postgres14.0]
redis: [4, 5, 6] redis: [6]
mongo: ["6.0"] mongo: ["6.0"]
if: > if: >
${{ github.event.workflow_run.event == 'pull_request' && ${{ github.event.workflow_run.event == 'pull_request' &&

View file

@ -924,7 +924,7 @@ jobs:
strategy: strategy:
matrix: matrix:
go: ["1.25"] go: ["1.25"]
db: [MySQL8.0, SQLite3] db: [MySQL8.0, SQLite3, Postgres14.0]
steps: steps:
- name: Checkout Kun - name: Checkout Kun
uses: actions/checkout@v4 uses: actions/checkout@v4
@ -1021,6 +1021,8 @@ jobs:
echo "YAO_DB_DRIVER=$DB_DRIVER" >> $GITHUB_ENV echo "YAO_DB_DRIVER=$DB_DRIVER" >> $GITHUB_ENV
if [ "$DB_DRIVER" = "mysql" ]; then if [ "$DB_DRIVER" = "mysql" ]; then
echo "YAO_DB_PRIMARY=$DB_USER:$PASSWORD@$DB_HOST" >> $GITHUB_ENV echo "YAO_DB_PRIMARY=$DB_USER:$PASSWORD@$DB_HOST" >> $GITHUB_ENV
elif [ "$DB_DRIVER" = "postgres" ]; then
echo "YAO_DB_PRIMARY=postgres://$DB_USER:$PASSWORD@$DB_HOST" >> $GITHUB_ENV
else else
echo "YAO_DB_PRIMARY=${{ github.WORKSPACE }}/../app/db/yao.db" >> $GITHUB_ENV echo "YAO_DB_PRIMARY=${{ github.WORKSPACE }}/../app/db/yao.db" >> $GITHUB_ENV
mkdir -p ${{ github.WORKSPACE }}/../app/db mkdir -p ${{ github.WORKSPACE }}/../app/db
@ -1066,8 +1068,8 @@ jobs:
strategy: strategy:
matrix: matrix:
go: ["1.25"] go: ["1.25"]
db: [MySQL8.0, SQLite3] db: [MySQL8.0, SQLite3, Postgres14.0]
redis: [4, 5, 6] redis: [6]
mongo: ["6.0"] mongo: ["6.0"]
steps: steps:
- name: Checkout Kun - name: Checkout Kun

1
.gitignore vendored
View file

@ -82,3 +82,4 @@ agent/robot/ROBOT-IM-INTEGRATION-IMPROVEMENT.md
agent/robot/ROBOT-CACHE-IMPROVEMENT.md agent/robot/ROBOT-CACHE-IMPROVEMENT.md
sandbox/v2/PID-KILL-UPGRADE.md sandbox/v2/PID-KILL-UPGRADE.md
sandbox/v2/*.md sandbox/v2/*.md
POSTGRESQL_COMPAT.md

View file

@ -41,6 +41,12 @@ unit-test:
if grep -q "^--- FAIL" tmp.out; then \ if grep -q "^--- FAIL" tmp.out; then \
rm tmp.out; \ rm tmp.out; \
exit 1; \ exit 1; \
elif grep -q "^FAIL" tmp.out; then \
rm tmp.out; \
exit 1; \
elif grep -q "^panic:" tmp.out; then \
rm tmp.out; \
exit 1; \
elif grep -q "build failed" tmp.out; then \ elif grep -q "build failed" tmp.out; then \
rm tmp.out; \ rm tmp.out; \
exit 1; \ exit 1; \
@ -67,6 +73,12 @@ unit-test-core:
if grep -q "^--- FAIL" tmp.out; then \ if grep -q "^--- FAIL" tmp.out; then \
rm tmp.out; \ rm tmp.out; \
exit 1; \ exit 1; \
elif grep -q "^FAIL" tmp.out; then \
rm tmp.out; \
exit 1; \
elif grep -q "^panic:" tmp.out; then \
rm tmp.out; \
exit 1; \
elif grep -q "build failed" tmp.out; then \ elif grep -q "build failed" tmp.out; then \
rm tmp.out; \ rm tmp.out; \
exit 1; \ exit 1; \
@ -359,6 +371,12 @@ unit-test-grpc:
if grep -q "^--- FAIL" tmp.out; then \ if grep -q "^--- FAIL" tmp.out; then \
rm tmp.out; \ rm tmp.out; \
exit 1; \ exit 1; \
elif grep -q "^FAIL" tmp.out; then \
rm tmp.out; \
exit 1; \
elif grep -q "^panic:" tmp.out; then \
rm tmp.out; \
exit 1; \
elif grep -q "build failed" tmp.out; then \ elif grep -q "build failed" tmp.out; then \
rm tmp.out; \ rm tmp.out; \
exit 1; \ exit 1; \

View file

@ -12,6 +12,7 @@ import (
"github.com/yaoapp/yao/agent/search" "github.com/yaoapp/yao/agent/search"
searchTypes "github.com/yaoapp/yao/agent/search/types" searchTypes "github.com/yaoapp/yao/agent/search/types"
store "github.com/yaoapp/yao/agent/store/types" store "github.com/yaoapp/yao/agent/store/types"
"github.com/yaoapp/yao/openapi/utils"
sui "github.com/yaoapp/yao/sui/core" sui "github.com/yaoapp/yao/sui/core"
) )
@ -147,8 +148,8 @@ func (ast *Assistant) Map() map[string]interface{} {
"uses": ast.Uses, "uses": ast.Uses,
"search": ast.Search, "search": ast.Search,
"dependencies": ast.Dependencies, "dependencies": ast.Dependencies,
"created_at": store.ToMySQLTime(ast.CreatedAt), "created_at": utils.NanoToTime(ast.CreatedAt),
"updated_at": store.ToMySQLTime(ast.UpdatedAt), "updated_at": utils.NanoToTime(ast.UpdatedAt),
} }
} }

View file

@ -326,12 +326,28 @@ func ToAssistantModel(v interface{}) (*AssistantModel, error) {
model.CreatedAt = createdAt model.CreatedAt = createdAt
} else if createdAt, ok := data["created_at"].(float64); ok { } else if createdAt, ok := data["created_at"].(float64); ok {
model.CreatedAt = int64(createdAt) model.CreatedAt = int64(createdAt)
} else if createdAt, ok := data["created_at"].(time.Time); ok {
model.CreatedAt = createdAt.UnixNano()
} else if createdAt, ok := data["created_at"].(string); ok && createdAt != "" {
if ts, err := time.Parse(time.RFC3339Nano, createdAt); err == nil {
model.CreatedAt = ts.UnixNano()
} else if ts, err := time.Parse("2006-01-02 15:04:05", createdAt); err == nil {
model.CreatedAt = ts.UnixNano()
}
} }
if updatedAt, ok := data["updated_at"].(int64); ok { if updatedAt, ok := data["updated_at"].(int64); ok {
model.UpdatedAt = updatedAt model.UpdatedAt = updatedAt
} else if updatedAt, ok := data["updated_at"].(float64); ok { } else if updatedAt, ok := data["updated_at"].(float64); ok {
model.UpdatedAt = int64(updatedAt) model.UpdatedAt = int64(updatedAt)
} else if updatedAt, ok := data["updated_at"].(time.Time); ok {
model.UpdatedAt = updatedAt.UnixNano()
} else if updatedAt, ok := data["updated_at"].(string); ok && updatedAt != "" {
if ts, err := time.Parse(time.RFC3339Nano, updatedAt); err == nil {
model.UpdatedAt = ts.UnixNano()
} else if ts, err := time.Parse("2006-01-02 15:04:05", updatedAt); err == nil {
model.UpdatedAt = ts.UnixNano()
}
} }
// Tags (string array) // Tags (string array)

View file

@ -69,22 +69,20 @@ func (store *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error)
data["automated"] = assistant.Automated data["automated"] = assistant.Automated
data["disable_global_prompts"] = assistant.DisableGlobalPrompts data["disable_global_prompts"] = assistant.DisableGlobalPrompts
// Set timestamps // Use UTC time.Time so the DB driver serialises correctly for all dialects
now := time.Now().UnixNano() // (PostgreSQL timestamptz, MySQL datetime, SQLite text) with no TZ ambiguity.
now := time.Now().UTC()
if exists { if exists {
// Update: set updated_at, keep created_at unchanged
if assistant.UpdatedAt == 0 { if assistant.UpdatedAt == 0 {
data["updated_at"] = now data["updated_at"] = now
} else { } else {
data["updated_at"] = assistant.UpdatedAt data["updated_at"] = nanoToTime(assistant.UpdatedAt)
} }
// Don't modify created_at on update
} else { } else {
// Create: set created_at, updated_at is null
if assistant.CreatedAt == 0 { if assistant.CreatedAt == 0 {
data["created_at"] = now data["created_at"] = now
} else { } else {
data["created_at"] = assistant.CreatedAt data["created_at"] = nanoToTime(assistant.CreatedAt)
} }
data["updated_at"] = nil data["updated_at"] = nil
} }
@ -272,8 +270,8 @@ func (store *Xun) UpdateAssistant(assistantID string, updates map[string]interfa
} }
} }
// Always update updated_at timestamp // Always update updated_at timestamp (UTC time.Time for dialect portability)
data["updated_at"] = types.ToMySQLTime(time.Now().UnixNano()) data["updated_at"] = time.Now().UTC()
if len(data) == 0 { if len(data) == 0 {
return fmt.Errorf("no valid fields to update") return fmt.Errorf("no valid fields to update")
@ -319,13 +317,11 @@ func (store *Xun) GetAssistants(filter types.AssistantFilter, locale ...string)
if len(filter.Tags) > 0 { if len(filter.Tags) > 0 {
qb.Where(func(qb query.Query) { qb.Where(func(qb query.Query) {
for i, tag := range filter.Tags { for i, tag := range filter.Tags {
// For each tag, we need to match it as part of a JSON array val := store.jsonContainsValue(fmt.Sprintf("%%\"%s\"%%", tag))
// This will match both single tag arrays ["tag1"] and multi-tag arrays ["tag1","tag2"]
pattern := fmt.Sprintf("%%\"%s\"%%", tag)
if i == 0 { if i == 0 {
qb.Where("tags", "like", pattern) qb.WhereJSONContains("tags", val)
} else { } else {
qb.OrWhere("tags", "like", pattern) qb.OrWhereJSONContains("tags", val)
} }
} }
}) })
@ -333,11 +329,13 @@ func (store *Xun) GetAssistants(filter types.AssistantFilter, locale ...string)
// Apply keyword filter if provided // Apply keyword filter if provided
if filter.Keywords != "" { if filter.Keywords != "" {
kw := fmt.Sprintf("%%%s%%", filter.Keywords)
qb.Where(func(qb query.Query) { qb.Where(func(qb query.Query) {
qb.Where("name", "like", fmt.Sprintf("%%%s%%", filter.Keywords)). qb.Where("name", "like", kw).
OrWhere("description", "like", fmt.Sprintf("%%%s%%", filter.Keywords)). OrWhere("description", "like", kw).
OrWhere("capabilities", "like", fmt.Sprintf("%%%s%%", filter.Keywords)). OrWhere("capabilities", "like", kw)
OrWhere("locales", "like", fmt.Sprintf("%%%s%%", filter.Keywords)) localeVal := store.jsonContainsValue(kw)
qb.OrWhereJSONContains("locales", localeVal)
}) })
} }
@ -382,16 +380,15 @@ func (store *Xun) GetAssistants(filter types.AssistantFilter, locale ...string)
} }
// Apply sandbox filter (true = has sandbox config, false = no sandbox config) // Apply sandbox filter (true = has sandbox config, false = no sandbox config)
// MySQL JSON columns distinguish between SQL NULL and JSON literal null. // DB JSON columns distinguish between SQL NULL and JSON literal null.
// CAST(sandbox AS CHAR) returns 'null' for JSON null and NULL for SQL NULL. // Dialect-specific: MySQL uses CAST(... AS CHAR), PG uses ::text, SQLite uses CAST(... AS TEXT).
if filter.Sandbox != nil { if filter.Sandbox != nil {
notNull, isNull := store.sandboxRawSQL()
if *filter.Sandbox { if *filter.Sandbox {
qb.WhereNotNull("sandbox"). qb.WhereNotNull("sandbox").WhereRaw(notNull)
WhereRaw("CAST(`sandbox` AS CHAR) <> 'null'")
} else { } else {
qb.Where(func(qb query.Query) { qb.Where(func(qb query.Query) {
qb.WhereNull("sandbox"). qb.WhereNull("sandbox").OrWhereRaw(isNull)
OrWhereRaw("CAST(`sandbox` AS CHAR) = 'null'")
}) })
} }
} }
@ -714,11 +711,11 @@ func (store *Xun) DeleteAssistants(filter types.AssistantFilter) (int64, error)
if len(filter.Tags) > 0 { if len(filter.Tags) > 0 {
qb.Where(func(qb query.Query) { qb.Where(func(qb query.Query) {
for i, tag := range filter.Tags { for i, tag := range filter.Tags {
pattern := fmt.Sprintf("%%\"%s\"%%", tag) val := store.jsonContainsValue(fmt.Sprintf("%%\"%s\"%%", tag))
if i == 0 { if i == 0 {
qb.Where("tags", "like", pattern) qb.WhereJSONContains("tags", val)
} else { } else {
qb.OrWhere("tags", "like", pattern) qb.OrWhereJSONContains("tags", val)
} }
} }
}) })
@ -726,9 +723,10 @@ func (store *Xun) DeleteAssistants(filter types.AssistantFilter) (int64, error)
// Apply keyword filter if provided // Apply keyword filter if provided
if filter.Keywords != "" { if filter.Keywords != "" {
kw := fmt.Sprintf("%%%s%%", filter.Keywords)
qb.Where(func(qb query.Query) { qb.Where(func(qb query.Query) {
qb.Where("name", "like", fmt.Sprintf("%%%s%%", filter.Keywords)). qb.Where("name", "like", kw).
OrWhere("description", "like", fmt.Sprintf("%%%s%%", filter.Keywords)) OrWhere("description", "like", kw)
}) })
} }
@ -906,3 +904,19 @@ func (store *Xun) translate(model *types.AssistantModel, assistantID string, loc
// in their original (English) form so that filter.tags round-trips // in their original (English) form so that filter.tags round-trips
// correctly through the LIKE query on the DB column. // correctly through the LIKE query on the DB column.
} }
// sandboxRawSQL returns dialect-specific raw SQL fragments for sandbox JSON null detection.
// Returns (notNullExpr, isNullExpr) for filtering sandbox field.
// WhereRaw is used here because this is a JSON literal `null` comparison, not a JSON array
// contains query. Each dialect requires different casting to compare the JSON value as text.
// No bind parameters are needed (pure string comparison), so no placeholder issues.
func (store *Xun) sandboxRawSQL() (string, string) {
switch store.getDriver() {
case "postgres":
return `"sandbox"::text <> 'null'`, `"sandbox"::text = 'null'`
case "sqlite3":
return `CAST(sandbox AS TEXT) <> 'null'`, `CAST(sandbox AS TEXT) = 'null'`
default:
return "CAST(`sandbox` AS CHAR) <> 'null'", "CAST(`sandbox` AS CHAR) = 'null'"
}
}

View file

@ -6,6 +6,7 @@ import (
"time" "time"
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/yao/openapi/utils"
) )
// isNil checks whether a value is truly nil, handling the Go typed-nil-in-interface pitfall. // isNil checks whether a value is truly nil, handling the Go typed-nil-in-interface pitfall.
@ -49,61 +50,20 @@ func getString(data map[string]interface{}, key string) string {
} }
func getBool(data map[string]interface{}, key string) bool { func getBool(data map[string]interface{}, key string) bool {
switch v := data[key].(type) { return utils.ToBool(data[key])
case bool:
return v
case int64:
return v != 0
case int:
return v != 0
case float64:
return v != 0
}
return false
} }
func getInt(data map[string]interface{}, key string) int { func getInt(data map[string]interface{}, key string) int {
switch v := data[key].(type) { return utils.ToInt(data[key])
case int:
return v
case int64:
return int(v)
case float64:
return int(v)
}
return 0
} }
func getInt64(data map[string]interface{}, key string) int64 { func getInt64(data map[string]interface{}, key string) int64 {
switch v := data[key].(type) { v := data[key]
case int64: if t, ok := v.(time.Time); ok {
return v return t.UnixNano()
case int:
return int64(v)
case float64:
return int64(v)
case string:
// Handle string representation of numbers (common with MySQL BIGINT)
var result int64
if _, err := fmt.Sscanf(v, "%d", &result); err == nil {
return result
}
case time.Time:
// Handle time.Time from database
return v.UnixNano()
} }
return 0 return utils.ToInt64(v)
} }
// toMySQLTime converts UnixNano timestamp to MySQL BIGINT format func nanoToTime(ns int64) time.Time { return utils.NanoToTime(ns) }
func toMySQLTime(unixNano int64) int64 { func timeToNano(t time.Time) int64 { return utils.TimeToNano(t) }
if unixNano == 0 {
return 0
}
return unixNano
}
// fromMySQLTime converts MySQL BIGINT timestamp to UnixNano
func fromMySQLTime(mysqlTime int64) int64 {
return mysqlTime
}

View file

@ -2,6 +2,7 @@ package xun
import ( import (
"fmt" "fmt"
"strings"
"time" "time"
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
@ -187,6 +188,26 @@ func (store *Xun) parseJSONFields(data map[string]interface{}, fields []string)
} }
} }
// getDriver returns the database driver name for dialect-aware SQL.
// Defaults to "mysql" if the driver cannot be determined.
func (store *Xun) getDriver() string {
if store.query != nil {
if driver, err := store.query.Driver(); err == nil {
return driver
}
}
return "mysql"
}
// jsonContainsValue formats a value for WhereJSONContains.
// PG/MySQL need JSON string (e.g. `"tag"`), SQLite needs LIKE pattern (e.g. `%"tag"%`).
func (store *Xun) jsonContainsValue(value string) string {
if store.getDriver() == "sqlite3" {
return value
}
return strings.TrimSuffix(strings.TrimPrefix(value, "%"), "%")
}
// GenerateAssistantID generates a random-looking 6-digit ID // GenerateAssistantID generates a random-looking 6-digit ID
func (store *Xun) GenerateAssistantID() (string, error) { func (store *Xun) GenerateAssistantID() (string, error) {
maxAttempts := 10 // Maximum number of attempts to generate a unique ID maxAttempts := 10 // Maximum number of attempts to generate a unique ID

View file

@ -0,0 +1,129 @@
package xun
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/yaoapp/xun/capsule"
"github.com/yaoapp/yao/agent/store/types"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/test"
)
func TestGetDriver(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
store, err := NewXun(typeSetting("default"))
assert.NoError(t, err)
xunStore := store.(*Xun)
driver := xunStore.getDriver()
assert.Contains(t, []string{"mysql", "postgres", "sqlite3"}, driver)
cfg := config.Conf
if cfg.DB.Driver != "" {
assert.Equal(t, cfg.DB.Driver, driver)
}
}
func TestSandboxRawSQL(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
store, err := NewXun(typeSetting("default"))
assert.NoError(t, err)
xunStore := store.(*Xun)
notNull, isNull := xunStore.sandboxRawSQL()
assert.NotEmpty(t, notNull)
assert.NotEmpty(t, isNull)
assert.Contains(t, notNull, "null")
assert.Contains(t, isNull, "null")
driver := xunStore.getDriver()
switch driver {
case "postgres":
assert.Contains(t, notNull, `"sandbox"::text`)
assert.Contains(t, isNull, `"sandbox"::text`)
case "sqlite3":
assert.Contains(t, notNull, "CAST(sandbox AS TEXT)")
assert.Contains(t, isNull, "CAST(sandbox AS TEXT)")
default:
assert.Contains(t, notNull, "CAST(`sandbox` AS CHAR)")
assert.Contains(t, isNull, "CAST(`sandbox` AS CHAR)")
}
}
func TestGetDriverFallback(t *testing.T) {
store := &Xun{}
driver := store.getDriver()
assert.Equal(t, "mysql", driver)
}
func typeSetting(connector string) types.Setting {
return types.Setting{Connector: connector}
}
// TestSandboxRawSQLAllDialects verifies SQL fragments for all three dialects
// without requiring a live database connection.
func TestSandboxRawSQLAllDialects(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
s, err := NewXun(typeSetting("default"))
assert.NoError(t, err)
xunStore := s.(*Xun)
// Verify the current driver produces valid SQL fragments
notNull, isNull := xunStore.sandboxRawSQL()
assert.Contains(t, notNull, "<>")
assert.Contains(t, isNull, "=")
}
func TestJsonContainsValue(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
store, err := NewXun(typeSetting("default"))
assert.NoError(t, err)
xunStore := store.(*Xun)
driver := xunStore.getDriver()
val := xunStore.jsonContainsValue(`%"admin"%`)
switch driver {
case "sqlite3":
assert.Equal(t, `%"admin"%`, val, "SQLite keeps LIKE pattern as-is")
default:
assert.Equal(t, `"admin"`, val, "PG/MySQL strips % wrappers for JSON value")
}
}
func TestJsonContainsValueFallback(t *testing.T) {
store := &Xun{}
val := store.jsonContainsValue(`%"test"%`)
assert.Equal(t, `"test"`, val, "Default (mysql) strips % wrappers")
}
func TestNanoToTime(t *testing.T) {
assert.True(t, nanoToTime(0).IsZero(), "zero input returns zero time")
ns := int64(1609459200000000000) // 2021-01-01 00:00:00 UTC
got := nanoToTime(ns)
assert.Equal(t, 2021, got.Year())
assert.Equal(t, time.January, got.Month())
assert.Equal(t, 1, got.Day())
assert.Equal(t, time.UTC, got.Location(), "must be UTC")
}
func TestTimeToNano(t *testing.T) {
assert.Equal(t, int64(0), timeToNano(time.Time{}), "zero time returns 0")
ts := time.Date(2021, 1, 1, 0, 0, 0, 0, time.UTC)
assert.Equal(t, int64(1609459200000000000), timeToNano(ts))
}
func init() {
_ = capsule.Global
}

File diff suppressed because one or more lines are too long

View file

@ -31,6 +31,7 @@ import (
"github.com/yaoapp/yao/flow" "github.com/yaoapp/yao/flow"
"github.com/yaoapp/yao/fs" "github.com/yaoapp/yao/fs"
"github.com/yaoapp/yao/i18n" "github.com/yaoapp/yao/i18n"
"github.com/yaoapp/yao/job"
"github.com/yaoapp/yao/kb" "github.com/yaoapp/yao/kb"
"github.com/yaoapp/yao/mcp" "github.com/yaoapp/yao/mcp"
"github.com/yaoapp/yao/messenger" "github.com/yaoapp/yao/messenger"
@ -477,6 +478,10 @@ func Unload() (err error) {
// Stop Runtime // Stop Runtime
err = runtime.Stop() err = runtime.Stop()
// Stop Job health checker and data cleaner before closing DB
job.StopHealthChecker()
job.StopDataCleaner()
// Close DB // Close DB
err = share.DBClose() err = share.DBClose()

View file

@ -1204,7 +1204,7 @@ func mapToStruct(m maps.MapStr, v interface{}) error {
case float64: case float64:
cleanMap[key] = val != 0 cleanMap[key] = val != 0
case string: case string:
cleanMap[key] = val == "true" || val == "1" cleanMap[key] = val == "true" || val == "1" || val == "t"
default: default:
cleanMap[key] = value cleanMap[key] = value
} }
@ -1213,9 +1213,12 @@ func mapToStruct(m maps.MapStr, v interface{}) error {
if str, ok := value.(string); ok && str != "" { if str, ok := value.(string); ok && str != "" {
// Try multiple time formats // Try multiple time formats
formats := []string{ formats := []string{
"2006-01-02 15:04:05", // MySQL format "2006-01-02 15:04:05",
"2006-01-02T15:04:05Z07:00", // RFC3339 "2006-01-02 15:04:05.999999",
"2006-01-02T15:04:05.999999999Z07:00", // RFC3339 with nanoseconds "2006-01-02 15:04:05.999999-07",
"2006-01-02 15:04:05.999999+00",
"2006-01-02T15:04:05Z07:00",
"2006-01-02T15:04:05.999999999Z07:00",
time.RFC3339, time.RFC3339,
time.RFC3339Nano, time.RFC3339Nano,
} }

97
job/data_internal_test.go Normal file
View file

@ -0,0 +1,97 @@
package job
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/yaoapp/kun/maps"
)
func TestMapToStructBooleanFields(t *testing.T) {
tests := []struct {
name string
key string
value interface{}
expected bool
}{
{"bool true", "enabled", true, true},
{"bool false", "enabled", false, false},
{"string true", "enabled", "true", true},
{"string 1", "enabled", "1", true},
{"string t (PG)", "enabled", "t", true},
{"string false", "enabled", "false", false},
{"string 0", "enabled", "0", false},
{"string f (PG)", "enabled", "f", false},
{"int 1", "enabled", int(1), true},
{"int 0", "enabled", int(0), false},
{"int64 1", "enabled", int64(1), true},
{"int64 0", "enabled", int64(0), false},
{"float64 1", "enabled", float64(1), true},
{"float64 0", "enabled", float64(0), false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
m := maps.MapStr{tt.key: tt.value}
job := &Job{}
err := mapToStruct(m, job)
assert.NoError(t, err)
assert.Equal(t, tt.expected, job.Enabled, "field %s with value %v", tt.key, tt.value)
})
}
t.Run("system field", func(t *testing.T) {
m := maps.MapStr{"system": "t"}
job := &Job{}
err := mapToStruct(m, job)
assert.NoError(t, err)
assert.True(t, job.System)
})
t.Run("readonly field", func(t *testing.T) {
m := maps.MapStr{"readonly": "1"}
job := &Job{}
err := mapToStruct(m, job)
assert.NoError(t, err)
assert.True(t, job.Readonly)
})
}
func TestMapToStructTimeFields(t *testing.T) {
refTime := time.Date(2024, 6, 15, 10, 30, 45, 0, time.UTC)
tests := []struct {
name string
input string
parsed bool
}{
{"MySQL format", "2024-06-15 10:30:45", true},
{"PG fractional", "2024-06-15 10:30:45.123456", true},
{"PG with +00", "2024-06-15 10:30:45.123456+00", true},
{"PG with -07", "2024-06-15 03:30:45.123456-07", true},
{"RFC3339", "2024-06-15T10:30:45Z", true},
{"RFC3339Nano", "2024-06-15T10:30:45.123456789Z", true},
{"empty string", "", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
m := maps.MapStr{
"name": "test-job",
"job_id": "j-123",
"created_at": tt.input,
}
job := &Job{}
err := mapToStruct(m, job)
if tt.parsed {
assert.NoError(t, err)
assert.False(t, job.CreatedAt.IsZero(), "expected non-zero time for input: %s", tt.input)
assert.Equal(t, refTime.Unix(), job.CreatedAt.Unix(),
"expected %v, got %v for input: %s", refTime, job.CreatedAt, tt.input)
} else {
assert.True(t, job.CreatedAt.IsZero())
}
})
}
}

View file

@ -4,6 +4,7 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"sync/atomic"
"time" "time"
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
@ -16,6 +17,8 @@ type HealthChecker struct {
interval time.Duration interval time.Duration
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
done chan struct{}
started atomic.Bool
} }
var globalHealthChecker *HealthChecker var globalHealthChecker *HealthChecker
@ -27,11 +30,20 @@ func NewHealthChecker(interval time.Duration) *HealthChecker {
interval: interval, interval: interval,
ctx: ctx, ctx: ctx,
cancel: cancel, cancel: cancel,
done: make(chan struct{}),
} }
} }
// Start starts the health check goroutine // Start starts the health check goroutine
func (hc *HealthChecker) Start() { func (hc *HealthChecker) Start() {
hc.started.Store(true)
defer close(hc.done)
defer func() {
if r := recover(); r != nil {
log.Error("Health checker recovered from panic: %v", r)
}
}()
ticker := time.NewTicker(hc.interval) ticker := time.NewTicker(hc.interval)
defer ticker.Stop() defer ticker.Stop()
@ -40,7 +52,7 @@ func (hc *HealthChecker) Start() {
for { for {
select { select {
case <-ticker.C: case <-ticker.C:
if err := hc.performHealthCheck(); err != nil { if err := hc.safePerformHealthCheck(); err != nil {
log.Error("Health check failed: %v", err) log.Error("Health check failed: %v", err)
} }
case <-hc.ctx.Done(): case <-hc.ctx.Done():
@ -50,11 +62,28 @@ func (hc *HealthChecker) Start() {
} }
} }
// Stop stops the health checker // safePerformHealthCheck wraps performHealthCheck with panic recovery
func (hc *HealthChecker) safePerformHealthCheck() (err error) {
defer func() {
if r := recover(); r != nil {
err = fmt.Errorf("health check panic: %v", r)
}
}()
return hc.performHealthCheck()
}
// Stop stops the health checker and waits for the goroutine to exit
func (hc *HealthChecker) Stop() { func (hc *HealthChecker) Stop() {
if hc.cancel != nil { if hc.cancel != nil {
hc.cancel() hc.cancel()
} }
if hc.started.Load() {
select {
case <-hc.done:
case <-time.After(5 * time.Second):
log.Error("Health checker stop timed out")
}
}
} }
// performHealthCheck performs health check // performHealthCheck performs health check
@ -261,6 +290,8 @@ func GetHealthChecker() *HealthChecker {
type DataCleaner struct { type DataCleaner struct {
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
done chan struct{}
started atomic.Bool
retentionDays int retentionDays int
lastCleanupTime time.Time lastCleanupTime time.Time
} }
@ -273,14 +304,22 @@ func NewDataCleaner(retentionDays int) *DataCleaner {
return &DataCleaner{ return &DataCleaner{
ctx: ctx, ctx: ctx,
cancel: cancel, cancel: cancel,
done: make(chan struct{}),
retentionDays: retentionDays, retentionDays: retentionDays,
lastCleanupTime: time.Now(), // Initialize to avoid immediate cleanup on startup lastCleanupTime: time.Now(),
} }
} }
// Start starts the daily data cleanup routine // Start starts the daily data cleanup routine
func (dc *DataCleaner) Start() { func (dc *DataCleaner) Start() {
// Check every hour if daily cleanup is needed dc.started.Store(true)
defer close(dc.done)
defer func() {
if r := recover(); r != nil {
log.Error("Data cleaner recovered from panic: %v", r)
}
}()
ticker := time.NewTicker(1 * time.Hour) ticker := time.NewTicker(1 * time.Hour)
defer ticker.Stop() defer ticker.Stop()
@ -290,7 +329,7 @@ func (dc *DataCleaner) Start() {
select { select {
case <-ticker.C: case <-ticker.C:
if dc.shouldRunCleanup() { if dc.shouldRunCleanup() {
if err := dc.performCleanup(); err != nil { if err := dc.safePerformCleanup(); err != nil {
log.Error("Data cleanup failed: %v", err) log.Error("Data cleanup failed: %v", err)
} else { } else {
dc.lastCleanupTime = time.Now() dc.lastCleanupTime = time.Now()
@ -303,11 +342,28 @@ func (dc *DataCleaner) Start() {
} }
} }
// Stop stops the data cleaner // safePerformCleanup wraps performCleanup with panic recovery
func (dc *DataCleaner) safePerformCleanup() (err error) {
defer func() {
if r := recover(); r != nil {
err = fmt.Errorf("data cleanup panic: %v", r)
}
}()
return dc.performCleanup()
}
// Stop stops the data cleaner and waits for the goroutine to exit
func (dc *DataCleaner) Stop() { func (dc *DataCleaner) Stop() {
if dc.cancel != nil { if dc.cancel != nil {
dc.cancel() dc.cancel()
} }
if dc.started.Load() {
select {
case <-dc.done:
case <-time.After(5 * time.Second):
log.Error("Data cleaner stop timed out")
}
}
} }
// shouldRunCleanup checks if cleanup should run (once per day) // shouldRunCleanup checks if cleanup should run (once per day)
@ -315,28 +371,30 @@ func (dc *DataCleaner) shouldRunCleanup() bool {
return time.Since(dc.lastCleanupTime) >= 24*time.Hour return time.Since(dc.lastCleanupTime) >= 24*time.Hour
} }
// performCleanup performs the actual data cleanup // performCleanup performs the actual data cleanup.
// Deletion order: logs -> executions -> jobs (children first) to avoid
// foreign key violations on databases that enforce referential integrity (e.g. PostgreSQL).
func (dc *DataCleaner) performCleanup() error { func (dc *DataCleaner) performCleanup() error {
log.Info("Starting daily data cleanup...") log.Info("Starting daily data cleanup...")
cutoffTime := time.Now().AddDate(0, 0, -dc.retentionDays) cutoffTime := time.Now().AddDate(0, 0, -dc.retentionDays)
// Clean up jobs (excluding running jobs) // Clean up logs first (leaf records)
deletedJobs, err := dc.cleanupJobs(cutoffTime) deletedLogs, err := dc.cleanupLogs(cutoffTime)
if err != nil { if err != nil {
return fmt.Errorf("failed to cleanup jobs: %w", err) return fmt.Errorf("failed to cleanup logs: %w", err)
} }
// Clean up executions // Clean up executions (reference jobs)
deletedExecutions, err := dc.cleanupExecutions(cutoffTime) deletedExecutions, err := dc.cleanupExecutions(cutoffTime)
if err != nil { if err != nil {
return fmt.Errorf("failed to cleanup executions: %w", err) return fmt.Errorf("failed to cleanup executions: %w", err)
} }
// Clean up logs // Clean up jobs last (parent records)
deletedLogs, err := dc.cleanupLogs(cutoffTime) deletedJobs, err := dc.cleanupJobs(cutoffTime)
if err != nil { if err != nil {
return fmt.Errorf("failed to cleanup logs: %w", err) return fmt.Errorf("failed to cleanup jobs: %w", err)
} }
log.Info("Data cleanup completed: %d jobs, %d executions, %d logs deleted", log.Info("Data cleanup completed: %d jobs, %d executions, %d logs deleted",

View file

@ -1,7 +1,7 @@
package types package types
import ( import (
"time" "github.com/yaoapp/yao/openapi/utils"
) )
// Map converts the OIDCUserInfo to a map[string]interface{}, excluding empty values // Map converts the OIDCUserInfo to a map[string]interface{}, excluding empty values
@ -67,7 +67,7 @@ func (user OIDCUserInfo) Map() map[string]interface{} {
} }
// Convert and add UpdatedAt if not nil // Convert and add UpdatedAt if not nil
if converted := unixToMySQL(user.UpdatedAt); converted != nil { if converted := utils.UnixToDBTimestamp(user.UpdatedAt); converted != nil {
result["updated_at"] = converted result["updated_at"] = converted
} }
@ -132,7 +132,7 @@ func (user OIDCUserInfo) Map() map[string]interface{} {
if user.YaoTeam.Description != "" { if user.YaoTeam.Description != "" {
teamMap["description"] = user.YaoTeam.Description teamMap["description"] = user.YaoTeam.Description
} }
if converted := unixToMySQL(user.YaoTeam.UpdatedAt); converted != nil { if converted := utils.UnixToDBTimestamp(user.YaoTeam.UpdatedAt); converted != nil {
teamMap["updated_at"] = converted teamMap["updated_at"] = converted
} }
if len(teamMap) > 0 { if len(teamMap) > 0 {
@ -264,7 +264,7 @@ func MakeOIDCUserInfo(user map[string]interface{}) *OIDCUserInfo {
// Updated_at field // Updated_at field
if updatedAt, ok := user["updated_at"]; ok { if updatedAt, ok := user["updated_at"]; ok {
if converted := toUnixTimestamp(updatedAt); converted != nil { if converted := utils.ToUnixTimestamp(updatedAt); converted != nil {
if unixTime, ok := converted.(int64); ok { if unixTime, ok := converted.(int64); ok {
userInfo.UpdatedAt = &unixTime userInfo.UpdatedAt = &unixTime
} }
@ -336,7 +336,7 @@ func MakeOIDCUserInfo(user map[string]interface{}) *OIDCUserInfo {
team.Description = description team.Description = description
} }
if updatedAt, ok := teamData["updated_at"]; ok { if updatedAt, ok := teamData["updated_at"]; ok {
if converted := toUnixTimestamp(updatedAt); converted != nil { if converted := utils.ToUnixTimestamp(updatedAt); converted != nil {
if unixTime, ok := converted.(int64); ok { if unixTime, ok := converted.(int64); ok {
team.UpdatedAt = &unixTime team.UpdatedAt = &unixTime
} }
@ -383,96 +383,3 @@ func MakeOIDCUserInfo(user map[string]interface{}) *OIDCUserInfo {
return userInfo return userInfo
} }
// unixToMySQL converts interface{} to MySQL DATETIME string
func unixToMySQL(val interface{}) interface{} {
if val == nil {
return nil
}
var unixTime int64
switch v := val.(type) {
case int64:
unixTime = v
case *int64:
if v == nil {
return nil
}
unixTime = *v
case int:
unixTime = int64(v)
case float64:
unixTime = int64(v)
default:
return nil
}
return time.Unix(unixTime, 0).UTC().Format("2006-01-02 15:04:05")
}
// mysqlToUnix converts interface{} to Unix timestamp
func mysqlToUnix(val interface{}) interface{} {
if val == nil {
return nil
}
var dateTime string
switch v := val.(type) {
case string:
dateTime = v
case *string:
if v == nil {
return nil
}
dateTime = *v
default:
return nil
}
if dateTime == "" {
return nil
}
// Try MySQL DATETIME format
if t, err := time.Parse("2006-01-02 15:04:05", dateTime); err == nil {
return t.Unix()
}
// Try ISO format as fallback
if t, err := time.Parse("2006-01-02T15:04:05Z", dateTime); err == nil {
return t.Unix()
}
return nil
}
// toUnixTimestamp converts any interface{} to Unix timestamp
func toUnixTimestamp(val interface{}) interface{} {
if val == nil {
return nil
}
switch v := val.(type) {
case int64:
return v
case *int64:
if v == nil {
return nil
}
return *v
case int:
return int64(v)
case float64:
return int64(v)
case string:
// Handle MySQL DATETIME or ISO format
return mysqlToUnix(v)
case *string:
if v == nil {
return nil
}
return mysqlToUnix(*v)
default:
return nil
}
}

View file

@ -863,6 +863,12 @@ func issueTokens(ctx context.Context, params *IssueTokensParams) (*LoginResponse
// prepareUserKBCollection prepares KB collection for user (called asynchronously after login) // prepareUserKBCollection prepares KB collection for user (called asynchronously after login)
func prepareUserKBCollection(userID, teamID, locale string) { func prepareUserKBCollection(userID, teamID, locale string) {
defer func() {
if r := recover(); r != nil {
log.Warn("prepareUserKBCollection recovered from panic: %v", r)
}
}()
// Get global KB setting // Get global KB setting
kbSetting := assistant.GetGlobalKBSetting() kbSetting := assistant.GetGlobalKBSetting()
if kbSetting == nil || kbSetting.Chat == nil { if kbSetting == nil || kbSetting.Chat == nil {

View file

@ -430,3 +430,46 @@ func buildProfileUpdateRequest(data map[string]interface{}) ProfileUpdateRequest
return req return req
} }
// GinProfileProviders returns the current user's linked OAuth accounts
func GinProfileProviders(c *gin.Context) {
authInfo := authorized.GetInfo(c)
if authInfo == nil || authInfo.UserID == "" {
response.RespondWithError(c, response.StatusUnauthorized, &response.ErrorResponse{
Code: response.ErrInvalidClient.Code,
ErrorDescription: "User not authenticated",
})
return
}
provider, err := getUserProvider()
if err != nil {
response.RespondWithError(c, response.StatusInternalServerError, &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to get user provider",
})
return
}
accounts, err := provider.GetUserOAuthAccounts(c.Request.Context(), authInfo.UserID)
if err != nil {
log.Error("Failed to get OAuth accounts for user %s: %v", authInfo.UserID, err)
response.RespondWithError(c, response.StatusInternalServerError, &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to retrieve linked accounts",
})
return
}
if accounts == nil {
accounts = []maps.MapStrAny{}
}
for i := range accounts {
if email, ok := accounts[i]["email"].(string); ok && email != "" {
accounts[i]["email"] = MaskEmail(email)
}
}
response.RespondWithSuccess(c, response.StatusOK, accounts)
}

View file

@ -250,8 +250,9 @@ func attachSubscription(group *gin.RouterGroup, oauth types.OAuth) {
// User profile management // User profile management
func attachProfile(group *gin.RouterGroup, oauth types.OAuth) { func attachProfile(group *gin.RouterGroup, oauth types.OAuth) {
group.GET("/profile", oauth.Guard, GinProfileGet) // Get user profile group.GET("/profile", oauth.Guard, GinProfileGet) // Get user profile
group.PUT("/profile", oauth.Guard, GinProfileUpdate) // Update user profile group.PUT("/profile", oauth.Guard, GinProfileUpdate) // Update user profile
group.GET("/profile/providers", oauth.Guard, GinProfileProviders) // Get linked OAuth providers
} }
// User management (CRUD) // User management (CRUD)

View file

@ -178,22 +178,12 @@ func ToTimeString(v interface{}) string {
} }
return val.Format(time.RFC3339) return val.Format(time.RFC3339)
case string: case string:
// Try to parse as RFC3339 first for _, format := range DBTimeFormats {
if t, err := time.Parse(time.RFC3339, val); err == nil {
return t.Format(time.RFC3339)
}
// Try to parse as other common formats
formats := []string{
"2006-01-02 15:04:05",
"2006-01-02T15:04:05Z",
"2006-01-02T15:04:05.000Z",
}
for _, format := range formats {
if t, err := time.Parse(format, val); err == nil { if t, err := time.Parse(format, val); err == nil {
return t.Format(time.RFC3339) return t.Format(time.RFC3339)
} }
} }
return val // Return as-is if can't parse return val
case int64: case int64:
// Assume unix timestamp // Assume unix timestamp
if val > 0 { if val > 0 {
@ -224,6 +214,131 @@ func GetTimeFormat(locale string) string {
} }
} }
// NanoToTime converts a UnixNano int64 to UTC time.Time.
// Returns zero time for zero input.
func NanoToTime(ns int64) time.Time {
if ns == 0 {
return time.Time{}
}
return time.Unix(ns/1e9, ns%1e9).UTC()
}
// TimeToNano converts time.Time to UnixNano int64.
// Returns 0 for zero time.
func TimeToNano(t time.Time) int64 {
if t.IsZero() {
return 0
}
return t.UnixNano()
}
// DBTimeFormats contains all time formats recognized by database drivers (MySQL, PostgreSQL, SQLite).
// Ordered from most specific to least specific for efficient parsing.
var DBTimeFormats = []string{
"2006-01-02 15:04:05",
"2006-01-02 15:04:05.999999",
"2006-01-02 15:04:05.999999-07",
"2006-01-02 15:04:05.999999+00",
"2006-01-02T15:04:05Z",
"2006-01-02T15:04:05",
"2006-01-02T15:04:05.000Z",
time.RFC3339,
time.RFC3339Nano,
}
// DBTimestampToUnix converts a database timestamp string (interface{}) to a Unix timestamp (int64).
// Supports MySQL DATETIME, PostgreSQL timestamptz (fractional seconds, timezone offsets),
// ISO 8601, RFC3339, and RFC3339Nano. Returns nil for nil, empty, or unparseable input.
func DBTimestampToUnix(val interface{}) interface{} {
if val == nil {
return nil
}
var dateTime string
switch v := val.(type) {
case string:
dateTime = v
case *string:
if v == nil {
return nil
}
dateTime = *v
default:
return nil
}
if dateTime == "" {
return nil
}
for _, format := range DBTimeFormats {
if t, err := time.Parse(format, dateTime); err == nil {
return t.Unix()
}
}
return nil
}
// UnixToDBTimestamp converts a Unix timestamp (int64, *int64, int, float64) to a standard
// DATETIME string "2006-01-02 15:04:05" in UTC. Compatible with MySQL, PostgreSQL, and SQLite.
// Returns nil for nil or unsupported types.
func UnixToDBTimestamp(val interface{}) interface{} {
if val == nil {
return nil
}
var unixTime int64
switch v := val.(type) {
case int64:
unixTime = v
case *int64:
if v == nil {
return nil
}
unixTime = *v
case int:
unixTime = int64(v)
case float64:
unixTime = int64(v)
default:
return nil
}
return time.Unix(unixTime, 0).UTC().Format("2006-01-02 15:04:05")
}
// ToUnixTimestamp converts any interface{} value to a Unix timestamp (int64).
// Handles numeric types directly and parses string/DB timestamp formats via DBTimestampToUnix.
// Returns nil for nil or unsupported types.
func ToUnixTimestamp(val interface{}) interface{} {
if val == nil {
return nil
}
switch v := val.(type) {
case int64:
return v
case *int64:
if v == nil {
return nil
}
return *v
case int:
return int64(v)
case float64:
return int64(v)
case string:
return DBTimestampToUnix(v)
case *string:
if v == nil {
return nil
}
return DBTimestampToUnix(*v)
default:
return nil
}
}
// FormatTimeWithLocale formats a time value (time.Time, *time.Time, or string) using the specified format // FormatTimeWithLocale formats a time value (time.Time, *time.Time, or string) using the specified format
// If the input is already a string, it will parse it first and then reformat it // If the input is already a string, it will parse it first and then reformat it
// Returns empty string if the value cannot be parsed // Returns empty string if the value cannot be parsed
@ -245,21 +360,13 @@ func FormatTimeWithLocale(v interface{}, targetFormat string) string {
return "" return ""
} }
case string: case string:
// Try parsing with common formats for _, format := range DBTimeFormats {
formats := []string{
"2006-01-02 15:04:05",
"2006-01-02T15:04:05Z",
"2006-01-02T15:04:05",
time.RFC3339,
}
for _, format := range formats {
t, err = time.Parse(format, val) t, err = time.Parse(format, val)
if err == nil { if err == nil {
break break
} }
} }
if err != nil { if err != nil {
// If all parsing attempts failed, return the original string
return val return val
} }
default: default:

View file

@ -0,0 +1,173 @@
package utils
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
)
func TestDBTimestampToUnix(t *testing.T) {
ref := time.Date(2024, 6, 15, 10, 30, 45, 0, time.UTC)
refUnix := ref.Unix()
tests := []struct {
name string
input interface{}
expect interface{}
}{
{"nil", nil, nil},
{"empty string", "", nil},
{"MySQL DATETIME", "2024-06-15 10:30:45", refUnix},
{"PG with fractional seconds", "2024-06-15 10:30:45.123456", refUnix},
{"PG with +00 offset", "2024-06-15 10:30:45.123456+00", refUnix},
{"PG with -07 offset", "2024-06-15 03:30:45.123456-07", refUnix},
{"ISO Z suffix", "2024-06-15T10:30:45Z", refUnix},
{"RFC3339", "2024-06-15T10:30:45+00:00", refUnix},
{"RFC3339Nano", "2024-06-15T10:30:45.123456789+00:00", refUnix},
{"*string nil", (*string)(nil), nil},
{"*string valid", strPtr("2024-06-15 10:30:45"), refUnix},
{"*string empty", strPtr(""), nil},
{"non-string type", 12345, nil},
{"unparseable string", "not-a-date", nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := DBTimestampToUnix(tt.input)
assert.Equal(t, tt.expect, result)
})
}
}
func TestUnixToDBTimestamp(t *testing.T) {
refUnix := int64(1718444445)
expected := time.Unix(refUnix, 0).UTC().Format("2006-01-02 15:04:05")
tests := []struct {
name string
input interface{}
expect interface{}
}{
{"nil", nil, nil},
{"int64", refUnix, expected},
{"*int64 valid", int64Ptr(refUnix), expected},
{"*int64 nil", (*int64)(nil), nil},
{"int", int(refUnix), expected},
{"float64", float64(refUnix), expected},
{"unknown type", "not-a-number", nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := UnixToDBTimestamp(tt.input)
assert.Equal(t, tt.expect, result)
})
}
}
func TestToUnixTimestamp(t *testing.T) {
refUnix := int64(1718444445)
tests := []struct {
name string
input interface{}
expect interface{}
}{
{"nil", nil, nil},
{"int64", refUnix, refUnix},
{"*int64 valid", int64Ptr(refUnix), refUnix},
{"*int64 nil", (*int64)(nil), nil},
{"int", int(refUnix), refUnix},
{"float64", float64(refUnix), refUnix},
{"string MySQL", "2024-06-15 10:00:45", time.Date(2024, 6, 15, 10, 0, 45, 0, time.UTC).Unix()},
{"string PG fractional", "2024-06-15 10:00:45.123456", time.Date(2024, 6, 15, 10, 0, 45, 0, time.UTC).Unix()},
{"string RFC3339", "2024-06-15T10:00:45Z", time.Date(2024, 6, 15, 10, 0, 45, 0, time.UTC).Unix()},
{"*string valid", strPtr("2024-06-15 10:00:45"), time.Date(2024, 6, 15, 10, 0, 45, 0, time.UTC).Unix()},
{"*string nil", (*string)(nil), nil},
{"string empty", "", nil},
{"unknown type", struct{}{}, nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := ToUnixTimestamp(tt.input)
assert.Equal(t, tt.expect, result)
})
}
}
func TestToTimeStringWithPGFormats(t *testing.T) {
tests := []struct {
name string
input interface{}
expect string
}{
{"PG fractional", "2024-06-15 10:30:45.123456", "2024-06-15T10:30:45Z"},
{"PG +00 offset", "2024-06-15 10:30:45.123456+00", "2024-06-15T10:30:45Z"},
{"PG -07 offset", "2024-06-15 03:30:45.123456-07", "2024-06-15T03:30:45-07:00"},
{"MySQL format", "2024-06-15 10:30:45", "2024-06-15T10:30:45Z"},
{"RFC3339Nano", "2024-06-15T10:30:45.123456789+00:00", "2024-06-15T10:30:45Z"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := ToTimeString(tt.input)
assert.Equal(t, tt.expect, result)
})
}
}
func TestNanoToTime(t *testing.T) {
t.Run("Zero", func(t *testing.T) {
assert.True(t, NanoToTime(0).IsZero())
})
t.Run("ValidTimestamp", func(t *testing.T) {
ns := int64(1609459200000000000) // 2021-01-01 00:00:00 UTC
got := NanoToTime(ns)
assert.Equal(t, 2021, got.Year())
assert.Equal(t, time.January, got.Month())
assert.Equal(t, 1, got.Day())
assert.Equal(t, 0, got.Hour())
assert.Equal(t, time.UTC, got.Location())
})
t.Run("PreservesNanoseconds", func(t *testing.T) {
ns := int64(1609459200123456789)
got := NanoToTime(ns)
assert.Equal(t, 123456789, got.Nanosecond())
})
t.Run("Negative", func(t *testing.T) {
got := NanoToTime(-1)
assert.False(t, got.IsZero())
assert.Equal(t, time.UTC, got.Location())
})
}
func TestTimeToNano(t *testing.T) {
t.Run("Zero", func(t *testing.T) {
assert.Equal(t, int64(0), TimeToNano(time.Time{}))
})
t.Run("ValidTime", func(t *testing.T) {
ts := time.Date(2021, 1, 1, 0, 0, 0, 0, time.UTC)
assert.Equal(t, int64(1609459200000000000), TimeToNano(ts))
})
t.Run("WithNanoseconds", func(t *testing.T) {
ts := time.Date(2021, 1, 1, 0, 0, 0, 123456789, time.UTC)
assert.Equal(t, int64(1609459200123456789), TimeToNano(ts))
})
t.Run("RoundTrip", func(t *testing.T) {
original := time.Date(2026, 3, 26, 15, 30, 45, 123456789, time.UTC)
ns := TimeToNano(original)
restored := NanoToTime(ns)
assert.True(t, original.Equal(restored))
})
}
func strPtr(s string) *string { return &s }
func int64Ptr(i int64) *int64 { return &i }

View file

@ -63,6 +63,14 @@ func withStaticFileServer(c *gin.Context) {
return return
} }
// Content negotiation: .md suffix or Accept: text/markdown
if strings.HasSuffix(c.Request.URL.Path, ".md") {
c.Set("content_type", "markdown")
c.Request.URL.Path = strings.TrimSuffix(c.Request.URL.Path, ".md")
} else if strings.Contains(c.GetHeader("Accept"), "text/markdown") {
c.Set("content_type", "markdown")
}
// Rewrite // Rewrite
for _, rewrite := range rewriteRules { for _, rewrite := range rewriteRules {
// log.Debug("Rewrite: %s => %s", c.Request.URL.Path, rewrite.Replacement) // log.Debug("Rewrite: %s => %s", c.Request.URL.Path, rewrite.Replacement)
@ -90,6 +98,18 @@ func withStaticFileServer(c *gin.Context) {
return return
} }
// Content negotiation: serve raw markdown instead of HTML
if ct, exists := c.Get("content_type"); exists && ct == "markdown" {
raw, contentType, code, err := r.RenderRaw("markdown")
if err != nil {
c.AbortWithStatusJSON(code, gin.H{"code": code, "message": err.Error()})
return
}
c.Data(code, contentType, []byte(raw))
c.Done()
return
}
html, code, err := r.Render() html, code, err := r.Render()
if err != nil { if err != nil {
if code == 301 || code == 302 { if code == 301 || code == 302 {

View file

@ -33,10 +33,13 @@ func TestStartStop(t *testing.T) {
} }
defer srv.Stop() defer srv.Stop()
<-srv.Event() port, err := srv.http.Port()
if err != nil {
t.Fatal(err)
}
// API Server // API Server
req := test.NewRequest(cfg.Port).Route("/api/__yao/app/setting") req := test.NewRequest(port).Route("/api/__yao/app/setting")
res, err := req.Get() res, err := req.Get()
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@ -49,7 +52,7 @@ func TestStartStop(t *testing.T) {
assert.True(t, len(data["name"].(string)) > 0) assert.True(t, len(data["name"].(string)) > 0)
// Public // Public
req = test.NewRequest(cfg.Port).Route("/") req = test.NewRequest(port).Route("/")
res, err = req.Get() res, err = req.Get()
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@ -58,7 +61,7 @@ func TestStartStop(t *testing.T) {
assert.Equal(t, "Hello World\n", res.Body()) assert.Equal(t, "Hello World\n", res.Body())
// XGEN // XGEN
req = test.NewRequest(cfg.Port).Route("/admin/") req = test.NewRequest(port).Route("/admin/")
res, err = req.Get() res, err = req.Get()
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)

View file

@ -12,6 +12,7 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/gou/application" "github.com/yaoapp/gou/application"
"github.com/yaoapp/gou/process"
"github.com/yaoapp/kun/exception" "github.com/yaoapp/kun/exception"
"github.com/yaoapp/kun/log" "github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/sui/core" "github.com/yaoapp/yao/sui/core"
@ -194,6 +195,122 @@ func (r *Request) Render() (string, int, error) {
return html, 200, nil return html, 200, nil
} }
// RenderRaw serves alternative content types (e.g. markdown) by invoking the
// handler declared in the page config, skipping data scripts and HTML rendering.
// "method" calls the page's backend.ts function; "process" calls a global Yao process.
func (r *Request) RenderRaw(kind string) (string, string, int, error) {
// Load or build cache (same as Render)
var c *core.Cache = nil
if !r.Request.DisableCache() {
c = core.GetCache(r.File)
}
if c == nil {
var status int
var err error
c, status, err = r.MakeCache()
if err != nil {
return "", "", status, err
}
}
// Guard
code, err := r.Guard(c)
if err != nil {
return "", "", code, err
}
// Parse config to find the handler
if c.Config == "" {
return "", "", 404, fmt.Errorf("page does not support %s output", kind)
}
var conf core.PageConfig
if err := jsoniter.UnmarshalFromString(c.Config, &conf); err != nil {
return "", "", 500, fmt.Errorf("config parse error: %s", err.Error())
}
// Resolve the handler from config by kind
var handler *core.PageProcess
switch kind {
case "markdown":
handler = conf.Markdown
}
if handler == nil || (handler.Method == "" && handler.Process == "") {
return "", "", 404, fmt.Errorf("page does not support %s output", kind)
}
// Resolve arguments
args := make([]interface{}, len(handler.In))
for i, expr := range handler.In {
args[i] = r.resolveArg(expr)
}
var result interface{}
if handler.Method != "" {
// Call backend.ts function via the page's compiled script
if c.Script == nil {
return "", "", 500, fmt.Errorf("page has no backend script")
}
r.Request.Script = c.Script
result, err = c.Script.Call(r.Request, handler.Method, args...)
if err != nil {
return "", "", 500, fmt.Errorf("backend script error: %s", err.Error())
}
} else {
// Fallback: call a global Yao process
p, err := process.Of(handler.Process, args...)
if err != nil {
return "", "", 500, fmt.Errorf("process error: %s", err.Error())
}
result, err = p.Exec()
if err != nil {
return "", "", 500, fmt.Errorf("process exec error: %s", err.Error())
}
}
content := ""
switch v := result.(type) {
case string:
content = v
case []byte:
content = string(v)
default:
return "", "", 500, fmt.Errorf("handler must return string, got %T", result)
}
contentType := "text/markdown; charset=utf-8"
return content, contentType, 200, nil
}
// resolveArg replaces $param.*, $query.* placeholders with actual request values.
func (r *Request) resolveArg(expr string) interface{} {
if strings.HasPrefix(expr, "$param.") {
key := expr[7:]
if val, ok := r.Request.Params[key]; ok {
return val
}
return ""
}
if strings.HasPrefix(expr, "$query.") {
key := expr[7:]
if r.Request.Query.Has(key) {
return r.Request.Query.Get(key)
}
return ""
}
if strings.HasPrefix(expr, "$header.") {
key := expr[8:]
if r.Request.Headers.Has(key) {
return r.Request.Headers.Get(key)
}
return ""
}
return expr
}
// MakeCache is the cache for the page API. // MakeCache is the cache for the page API.
func (r *Request) MakeCache() (*core.Cache, int, error) { func (r *Request) MakeCache() (*core.Cache, int, error) {

View file

@ -72,6 +72,7 @@ func (page *Page) ExportConfig() string {
"dataCache": page.Config.DataCache, "dataCache": page.Config.DataCache,
"api": page.Config.API, "api": page.Config.API,
"root": page.Root, "root": page.Root,
"markdown": page.Config.Markdown,
}) })
if err != nil { if err != nil {

View file

@ -375,15 +375,25 @@ type PageConfig struct {
// PageSetting is the struct for the page setting // PageSetting is the struct for the page setting
type PageSetting struct { type PageSetting struct {
Title string `json:"title,omitempty"` Title string `json:"title,omitempty"`
Guard string `json:"guard,omitempty"` Guard string `json:"guard,omitempty"`
CacheStore string `json:"cacheStore,omitempty"` CacheStore string `json:"cacheStore,omitempty"`
Cache int `json:"cache,omitempty"` Cache int `json:"cache,omitempty"`
Root string `json:"root,omitempty"` Root string `json:"root,omitempty"`
DataCache int `json:"dataCache,omitempty"` DataCache int `json:"dataCache,omitempty"`
Description string `json:"description,omitempty"` Description string `json:"description,omitempty"`
SEO *PageSEO `json:"seo,omitempty"` SEO *PageSEO `json:"seo,omitempty"`
API *PageAPI `json:"api,omitempty"` API *PageAPI `json:"api,omitempty"`
Markdown *PageProcess `json:"markdown,omitempty"`
}
// PageProcess binds a handler to a content negotiation output.
// "method" calls the page's own backend.ts function (preferred);
// "process" calls a global Yao process as fallback.
type PageProcess struct {
Method string `json:"method,omitempty"`
Process string `json:"process,omitempty"`
In []string `json:"in,omitempty"`
} }
// PageConfigRendered is the struct for the page config rendered // PageConfigRendered is the struct for the page config rendered

View file

@ -571,6 +571,8 @@ func dbconnect(t *testing.T, cfg config.Config) {
switch cfg.DB.Driver { switch cfg.DB.Driver {
case "sqlite3": case "sqlite3":
capsule.AddConn("primary", "sqlite3", cfg.DB.Primary[0]).SetAsGlobal() capsule.AddConn("primary", "sqlite3", cfg.DB.Primary[0]).SetAsGlobal()
case "postgres":
capsule.AddConn("primary", "postgres", cfg.DB.Primary[0]).SetAsGlobal()
default: default:
capsule.AddConn("primary", "mysql", cfg.DB.Primary[0]).SetAsGlobal() capsule.AddConn("primary", "mysql", cfg.DB.Primary[0]).SetAsGlobal()
} }

View file

@ -32,6 +32,7 @@ func TestProcessComponent(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
prepare(t) prepare(t)
clear(t)
testData(t) testData(t)
args := []interface{}{ args := []interface{}{

View file

@ -102,7 +102,7 @@
"comment": "User-specified complete file path", "comment": "User-specified complete file path",
"length": 1000, "length": 1000,
"nullable": true, "nullable": true,
"index": true "index": false
}, },
{ {
"name": "path", "name": "path",
@ -111,7 +111,7 @@
"comment": "Actual storage path for the file", "comment": "Actual storage path for the file",
"length": 1000, "length": 1000,
"nullable": false, "nullable": false,
"index": true "index": false
}, },
{ {
"name": "groups", "name": "groups",