仰望星辰工作室

clearlove2.1

clearlove2.1/ internal/plugin/api_db.go 11.0 KB · 422 行 原始文件
1// Host API:数据层(clv.db.* / clv.table)。
2//
3// 安全策略:
4// - 声明式权限 db 是使用数据能力的前提
5// - 写操作默认只允许本插件前缀(pl_<slug>_)的表;内核表只读
6// - 声明 db.admin 权限可放开全库写(需管理员在后台显式授权)
7// - 查询行数设上限,防止插件拉爆内存
8package plugin
9
10import (
11 "database/sql"
12 "errors"
13 "fmt"
14 "regexp"
15 "sort"
16 "strings"
17
18 "github.com/dop251/goja"
19
20 "clearlove/internal/database"
21)
22
23const maxQueryRows = 1000
24
25type sqlQuerier interface {
26 Query(query string, args ...any) (*sql.Rows, error)
27}
28
29type sqlExecer interface {
30 Exec(query string, args ...any) (sql.Result, error)
31}
32
33func (rt *jsRuntime) installDB(clv *goja.Object) error {
34 db := rt.vm.NewObject()
35
36 _ = db.Set("dialect", rt.jsFn(func(call goja.FunctionCall) (any, error) {
37 return database.Driver, nil
38 }))
39
40 _ = db.Set("table", rt.jsFn(func(call goja.FunctionCall) (any, error) {
41 args := argsOf(call)
42 if len(args) == 0 {
43 return nil, errors.New("clv.db.table(name) 需要一个参数")
44 }
45 return rt.plugin.TableName(strOf(args[0])), nil
46 }))
47
48 _ = db.Set("query", rt.jsFn(rt.dbQuery))
49 _ = db.Set("get", rt.jsFn(rt.dbGet))
50 _ = db.Set("exec", rt.jsFn(rt.dbExec))
51 _ = db.Set("tx", rt.jsFn(rt.dbTx))
52 _ = clv.Set("db", db)
53
54 // 建表:clv.table(name, columns, opts)
55 _ = clv.Set("table", rt.jsFn(rt.createTable))
56 return nil
57}
58
59// ---------- SQL 安全 ----------
60
61var (
62 reTableRef = regexp.MustCompile("(?i)\\b(?:into|update|from|table|join)\\s+[`\"']?([a-zA-Z0-9_]+)")
63 reSQLComment = regexp.MustCompile(`(?s)/\*.*?\*/|--[^\n]*`)
64 reDangerousSQL = regexp.MustCompile(`(?i)\b(pragma|attach|detach|vacuum)\b`)
65)
66
67// sqlAllowed 插件 SQL 管控:
68// - 声明 db.admin 权限的插件不受限
69// - 拒绝多语句(分号拆分逐段校验)与 PRAGMA/ATTACH/DETACH/VACUUM 等危险语句
70// - 读写一致:所有被引用的表必须落在本插件前缀(pl_<slug>_)内;
71// 读取内核表请走 clv.site(site.read 权限),防止仅声明 db 的插件拖库
72func (rt *jsRuntime) sqlAllowed(sqlText string) bool {
73 if rt.app != nil && rt.app.Plugin.HasPermission("db.admin") {
74 return true
75 }
76 // 去除注释再校验,防止 UPDATE/**/settings 这类拆词绕过
77 cleaned := reSQLComment.ReplaceAllString(sqlText, " ")
78 for _, seg := range strings.Split(cleaned, ";") {
79 seg = strings.TrimSpace(seg)
80 if seg == "" {
81 continue
82 }
83 if reDangerousSQL.MatchString(seg) {
84 return false
85 }
86 prefix := strings.ToLower(rt.plugin.TablePrefix())
87 for _, m := range reTableRef.FindAllStringSubmatch(seg, -1) {
88 if !strings.HasPrefix(strings.ToLower(m[1]), prefix) {
89 return false
90 }
91 }
92 }
93 return true
94}
95
96func (rt *jsRuntime) sqlDeniedErr(sqlText string) error {
97 return fmt.Errorf("插件 %s 只能访问自己的数据表(%s*);读取内核数据请用 clv.site,如需访问其它表请声明 db.admin 权限。SQL: %.100s",
98 rt.plugin.SlugOf(), rt.plugin.TablePrefix(), strings.TrimSpace(sqlText))
99}
100
101// ---------- db.query / db.get / db.exec ----------
102
103func (rt *jsRuntime) dbQuery(call goja.FunctionCall) (any, error) {
104 if err := rt.requirePerm("db"); err != nil {
105 return nil, err
106 }
107 args := argsOf(call)
108 if len(args) == 0 || strOf(args[0]) == "" {
109 return nil, errors.New("clv.db.query(sql, ...args) 需要 SQL")
110 }
111 sqlText := strOf(args[0])
112 if !rt.sqlAllowed(sqlText) {
113 return nil, rt.sqlDeniedErr(sqlText)
114 }
115 if database.DB == nil {
116 return nil, errors.New("数据库未连接")
117 }
118 return queryMaps(database.DB, sqlText, args[1:])
119}
120
121func (rt *jsRuntime) dbGet(call goja.FunctionCall) (any, error) {
122 if err := rt.requirePerm("db"); err != nil {
123 return nil, err
124 }
125 args := argsOf(call)
126 if len(args) == 0 || strOf(args[0]) == "" {
127 return nil, errors.New("clv.db.get(sql, ...args) 需要 SQL")
128 }
129 sqlText := strOf(args[0])
130 if !rt.sqlAllowed(sqlText) {
131 return nil, rt.sqlDeniedErr(sqlText)
132 }
133 if database.DB == nil {
134 return nil, errors.New("数据库未连接")
135 }
136 rows, err := queryMaps(database.DB, sqlText, args[1:])
137 if err != nil {
138 return nil, err
139 }
140 if len(rows) == 0 {
141 return nil, nil
142 }
143 return rows[0], nil
144}
145
146func (rt *jsRuntime) dbExec(call goja.FunctionCall) (any, error) {
147 if err := rt.requirePerm("db"); err != nil {
148 return nil, err
149 }
150 args := argsOf(call)
151 if len(args) == 0 || strOf(args[0]) == "" {
152 return nil, errors.New("clv.db.exec(sql, ...args) 需要 SQL")
153 }
154 sqlText := strOf(args[0])
155 if !rt.sqlAllowed(sqlText) {
156 return nil, rt.sqlDeniedErr(sqlText)
157 }
158 if database.DB == nil {
159 return nil, errors.New("数据库未连接")
160 }
161 return execResult(database.DB, sqlText, args[1:])
162}
163
164// dbTx 事务:clv.db.tx(function (tx) { tx.exec(...); tx.get(...); })
165func (rt *jsRuntime) dbTx(call goja.FunctionCall) (any, error) {
166 if err := rt.requirePerm("db"); err != nil {
167 return nil, err
168 }
169 fn, ok := goja.AssertFunction(call.Argument(0))
170 if !ok {
171 return nil, errors.New("clv.db.tx(fn) 需要传入函数")
172 }
173 if database.DB == nil {
174 return nil, errors.New("数据库未连接")
175 }
176 tx, err := database.DB.Begin()
177 if err != nil {
178 return nil, fmt.Errorf("开启事务失败: %w", err)
179 }
180 txObj := rt.vm.NewObject()
181 _ = txObj.Set("exec", rt.jsFn(func(c goja.FunctionCall) (any, error) {
182 a := argsOf(c)
183 if len(a) == 0 {
184 return nil, errors.New("tx.exec(sql, ...args) 需要 SQL")
185 }
186 s := strOf(a[0])
187 if !rt.sqlAllowed(s) {
188 return nil, rt.sqlDeniedErr(s)
189 }
190 return execResult(tx, s, a[1:])
191 }))
192 _ = txObj.Set("get", rt.jsFn(func(c goja.FunctionCall) (any, error) {
193 a := argsOf(c)
194 if len(a) == 0 {
195 return nil, errors.New("tx.get(sql, ...args) 需要 SQL")
196 }
197 s := strOf(a[0])
198 if !rt.sqlAllowed(s) {
199 return nil, rt.sqlDeniedErr(s)
200 }
201 rows, err := queryMaps(tx, s, a[1:])
202 if err != nil {
203 return nil, err
204 }
205 if len(rows) == 0 {
206 return nil, nil
207 }
208 return rows[0], nil
209 }))
210 _ = txObj.Set("query", rt.jsFn(func(c goja.FunctionCall) (any, error) {
211 a := argsOf(c)
212 if len(a) == 0 {
213 return nil, errors.New("tx.query(sql, ...args) 需要 SQL")
214 }
215 s := strOf(a[0])
216 if !rt.sqlAllowed(s) {
217 return nil, rt.sqlDeniedErr(s)
218 }
219 return queryMaps(tx, s, a[1:])
220 }))
221
222 v, err := fn(goja.Undefined(), txObj)
223 if err != nil {
224 _ = tx.Rollback()
225 return nil, err
226 }
227 if err := tx.Commit(); err != nil {
228 return nil, fmt.Errorf("提交事务失败: %w", err)
229 }
230 return exportJS(v), nil
231}
232
233// ---------- 建表 ----------
234
235// createTable clv.table(name, columns, opts) —— 幂等建表
236func (rt *jsRuntime) createTable(call goja.FunctionCall) (any, error) {
237 if err := rt.requirePerm("db"); err != nil {
238 return nil, err
239 }
240 if database.DB == nil {
241 return nil, errors.New("数据库未连接")
242 }
243 name := NormalizeSlug(strOf(call.Argument(0)))
244 if name == "" {
245 return nil, errors.New("clv.table(name, columns, opts) 需要表名(字母/数字/下划线)")
246 }
247 cols := mapOfAny(call.Argument(1))
248 if len(cols) == 0 {
249 return nil, errors.New("clv.table 需要列定义对象,例如 { user_id: 'int', day: 'string' }")
250 }
251 opts := mapOfAny(call.Argument(2))
252
253 table := rt.plugin.TableName(name)
254 ddl := buildCreateTable(database.Driver, table, cols, opts)
255 if _, err := database.DB.Exec(ddl); err != nil {
256 return nil, fmt.Errorf("建表 %s 失败: %w", table, err)
257 }
258 for _, ix := range buildIndexes(table, opts) {
259 _, _ = database.DB.Exec(ix) // MySQL 不支持 IF NOT EXISTS,重复创建的错误忽略
260 }
261 return table, nil
262}
263
264func sqlType(driver, t string) string {
265 t = strings.ToLower(strings.TrimSpace(t))
266 if driver == "mysql" {
267 switch t {
268 case "int":
269 return "INT"
270 case "bigint":
271 return "BIGINT"
272 case "float", "number":
273 return "DOUBLE"
274 case "bool":
275 return "TINYINT"
276 case "text":
277 return "MEDIUMTEXT"
278 case "json":
279 return "MEDIUMTEXT"
280 case "time":
281 return "VARCHAR(40)"
282 default:
283 return "VARCHAR(255)"
284 }
285 }
286 switch t {
287 case "int", "bigint", "bool":
288 return "INTEGER"
289 case "float", "number":
290 return "REAL"
291 case "time":
292 return "TEXT"
293 default:
294 return "TEXT"
295 }
296}
297
298func buildCreateTable(driver, table string, cols map[string]any, opts map[string]any) string {
299 var b strings.Builder
300 b.WriteString("CREATE TABLE IF NOT EXISTS " + table + " (")
301 if driver == "mysql" {
302 b.WriteString("id BIGINT UNSIGNED PRIMARY KEY AUTO_INCREMENT")
303 } else {
304 b.WriteString("id INTEGER PRIMARY KEY AUTOINCREMENT")
305 }
306 names := make([]string, 0, len(cols))
307 for k := range cols {
308 names = append(names, k)
309 }
310 sort.Strings(names)
311 for _, k := range names {
312 col := strings.ReplaceAll(NormalizeSlug(k), "-", "_")
313 if col == "" || col == "id" {
314 continue
315 }
316 b.WriteString(", " + col + " " + sqlType(driver, strOf(cols[k])))
317 }
318 for _, u := range listOfLists(opts["unique"]) {
319 b.WriteString(", UNIQUE(" + strings.Join(u, ",") + ")")
320 }
321 b.WriteString(")")
322 return b.String()
323}
324
325func buildIndexes(table string, opts map[string]any) []string {
326 var out []string
327 for i, cols := range listOfLists(opts["indexes"]) {
328 out = append(out, fmt.Sprintf("CREATE INDEX idx_%s_%d ON %s (%s)",
329 table, i, table, strings.Join(cols, ",")))
330 }
331 return out
332}
333
334// ---------- 通用查询辅助 ----------
335
336func queryMaps(q sqlQuerier, sqlText string, args []any) ([]map[string]any, error) {
337 rows, err := q.Query(sqlText, args...)
338 if err != nil {
339 return nil, err
340 }
341 defer rows.Close()
342 cols, err := rows.Columns()
343 if err != nil {
344 return nil, err
345 }
346 out := make([]map[string]any, 0, 16)
347 for rows.Next() {
348 if len(out) >= maxQueryRows {
349 return nil, fmt.Errorf("查询结果超过 %d 行上限,请加 LIMIT", maxQueryRows)
350 }
351 vals := make([]any, len(cols))
352 ptrs := make([]any, len(cols))
353 for i := range vals {
354 ptrs[i] = &vals[i]
355 }
356 if err := rows.Scan(ptrs...); err != nil {
357 return nil, err
358 }
359 m := make(map[string]any, len(cols))
360 for i, c := range cols {
361 m[c] = normalizeSQLValue(vals[i])
362 }
363 out = append(out, m)
364 }
365 return out, rows.Err()
366}
367
368func execResult(e sqlExecer, sqlText string, args []any) (map[string]any, error) {
369 res, err := e.Exec(sqlText, args...)
370 if err != nil {
371 return nil, err
372 }
373 lastID, _ := res.LastInsertId()
374 affected, _ := res.RowsAffected()
375 return map[string]any{"lastId": lastID, "rowsAffected": affected}, nil
376}
377
378// normalizeSQLValue 统一驱动差异([]byte -> string 等)
379func normalizeSQLValue(v any) any {
380 switch t := v.(type) {
381 case []byte:
382 return string(t)
383 default:
384 return v
385 }
386}
387
388// mapOfAny 取 goja 对象参数为 map
389func mapOfAny(v goja.Value) map[string]any {
390 if v == nil || goja.IsUndefined(v) || goja.IsNull(v) {
391 return nil
392 }
393 if m, ok := v.Export().(map[string]any); ok {
394 return m
395 }
396 return nil
397}
398
399// listOfLists 解析 [["a","b"],["c"]] 结构
400func listOfLists(v any) [][]string {
401 arr, ok := v.([]any)
402 if !ok {
403 return nil
404 }
405 out := make([][]string, 0, len(arr))
406 for _, item := range arr {
407 inner, ok := item.([]any)
408 if !ok {
409 continue
410 }
411 row := make([]string, 0, len(inner))
412 for _, c := range inner {
413 if col := strings.ReplaceAll(NormalizeSlug(strOf(c)), "-", "_"); col != "" {
414 row = append(row, col)
415 }
416 }
417 if len(row) > 0 {
418 out = append(out, row)
419 }
420 }
421 return out
422}