mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 23:56:36 +08:00
Compare commits
13 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3799729706 | |||
| 8628c35802 | |||
| acea5ea76c | |||
| 8652380da1 | |||
| 70f8dfeac1 | |||
| 37005a1954 | |||
| b55e056316 | |||
| b3b7f5e56d | |||
| d6ff6ea500 | |||
| 9ed875b7ef | |||
| 6387ce1816 | |||
| 219067a27c | |||
| 33678477aa |
@@ -7,18 +7,17 @@ services:
|
||||
driver: json-file
|
||||
options:
|
||||
max-size: "20m"
|
||||
max-file: "3"
|
||||
environment:
|
||||
DB_TYPE: ${DB_TYPE:-sqlite}
|
||||
DB_PATH: /app/data/gost.db
|
||||
DATABASE_URL: ${DATABASE_URL:-}
|
||||
JWT_SECRET: ${JWT_SECRET}
|
||||
LOG_DIR: /app/logs
|
||||
SERVER_ADDR: :6365
|
||||
TZ: Asia/Shanghai
|
||||
ports:
|
||||
- "${BACKEND_PORT}:6365"
|
||||
volumes:
|
||||
- backend_logs:/app/logs
|
||||
- sqlite_data:/app/data
|
||||
networks:
|
||||
- gost-network
|
||||
@@ -63,6 +62,7 @@ services:
|
||||
driver: json-file
|
||||
options:
|
||||
max-size: "20m"
|
||||
max-file: "3"
|
||||
ports:
|
||||
- "${FRONTEND_PORT}:80"
|
||||
depends_on:
|
||||
@@ -79,9 +79,6 @@ volumes:
|
||||
postgres_data:
|
||||
name: postgres_data
|
||||
driver: local
|
||||
backend_logs:
|
||||
name: backend_logs
|
||||
driver: local
|
||||
|
||||
|
||||
networks:
|
||||
|
||||
@@ -7,18 +7,17 @@ services:
|
||||
driver: json-file
|
||||
options:
|
||||
max-size: "20m"
|
||||
max-file: "3"
|
||||
environment:
|
||||
DB_TYPE: ${DB_TYPE:-sqlite}
|
||||
DB_PATH: /app/data/gost.db
|
||||
DATABASE_URL: ${DATABASE_URL:-}
|
||||
JWT_SECRET: ${JWT_SECRET}
|
||||
LOG_DIR: /app/logs
|
||||
SERVER_ADDR: :6365
|
||||
TZ: Asia/Shanghai
|
||||
ports:
|
||||
- "${BACKEND_PORT}:6365"
|
||||
volumes:
|
||||
- backend_logs:/app/logs
|
||||
- sqlite_data:/app/data
|
||||
networks:
|
||||
- gost-network
|
||||
@@ -63,6 +62,7 @@ services:
|
||||
driver: json-file
|
||||
options:
|
||||
max-size: "20m"
|
||||
max-file: "3"
|
||||
ports:
|
||||
- "${FRONTEND_PORT}:80"
|
||||
depends_on:
|
||||
@@ -79,9 +79,6 @@ volumes:
|
||||
postgres_data:
|
||||
name: postgres_data
|
||||
driver: local
|
||||
backend_logs:
|
||||
name: backend_logs
|
||||
driver: local
|
||||
|
||||
|
||||
networks:
|
||||
|
||||
@@ -889,11 +889,11 @@ func firstPortFromRange(portRange string) int {
|
||||
|
||||
func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
|
||||
rows, err := h.repo.DB().Query(`
|
||||
SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
|
||||
SELECT CAST(ct.chain_type AS INTEGER), COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
|
||||
FROM chain_tunnel ct
|
||||
LEFT JOIN node n ON n.id = ct.node_id
|
||||
WHERE ct.tunnel_id = ?
|
||||
ORDER BY ct.chain_type ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
|
||||
ORDER BY CAST(ct.chain_type AS INTEGER) ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
|
||||
`, tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -799,7 +799,7 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, 1, ?, ?, 'fifo', 0, ?)`,
|
||||
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, '1', ?, ?, 'fifo', 0, ?)`,
|
||||
tunnelID,
|
||||
share.NodeID,
|
||||
req.RemotePort,
|
||||
@@ -1000,14 +1000,19 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid target"))
|
||||
return
|
||||
}
|
||||
targetProtocol := defaultString(target.Protocol, protocol)
|
||||
connector := map[string]interface{}{
|
||||
"type": "relay",
|
||||
}
|
||||
if isTLSTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{"nodelay": true}
|
||||
}
|
||||
nodeItems = append(nodeItems, map[string]interface{}{
|
||||
"name": fmt.Sprintf("node_%d", i+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, target.Port)),
|
||||
"connector": map[string]interface{}{
|
||||
"type": "relay",
|
||||
},
|
||||
"name": fmt.Sprintf("node_%d", i+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, target.Port)),
|
||||
"connector": connector,
|
||||
"dialer": map[string]interface{}{
|
||||
"type": defaultString(target.Protocol, protocol),
|
||||
"type": targetProtocol,
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -1046,6 +1051,9 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
"type": protocol,
|
||||
},
|
||||
}
|
||||
if isTLSTunnelProtocol(protocol) {
|
||||
service["handler"].(map[string]interface{})["metadata"] = map[string]interface{}{"nodelay": true}
|
||||
}
|
||||
if req.Role == "middle" {
|
||||
service["handler"].(map[string]interface{})["chain"] = chainName
|
||||
}
|
||||
|
||||
@@ -306,11 +306,35 @@ func (h *Handler) userList(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Current int `json:"current"`
|
||||
Size int `json:"size"`
|
||||
Keyword string `json:"keyword"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil && err != io.EOF {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
users, err := h.repo.ListUsers()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
keyword := strings.ToLower(strings.TrimSpace(req.Keyword))
|
||||
if keyword != "" {
|
||||
filtered := make([]map[string]interface{}, 0, len(users))
|
||||
for _, item := range users {
|
||||
username := strings.ToLower(strings.TrimSpace(fmt.Sprint(item["user"])))
|
||||
displayName := strings.ToLower(strings.TrimSpace(fmt.Sprint(item["name"])))
|
||||
if strings.Contains(username, keyword) || strings.Contains(displayName, keyword) {
|
||||
filtered = append(filtered, item)
|
||||
}
|
||||
}
|
||||
users = filtered
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(users))
|
||||
}
|
||||
|
||||
|
||||
@@ -688,6 +688,9 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
runtimeState.TunnelID = id
|
||||
|
||||
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes)
|
||||
|
||||
var federationBindings []sqlite.FederationTunnelBinding
|
||||
var federationReleaseRefs []federationRuntimeReleaseRef
|
||||
if typeVal == 2 {
|
||||
@@ -700,7 +703,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
applyTunnelPortsToRequest(req, runtimeState)
|
||||
|
||||
_, err = tx.Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, updated_time=? WHERE id=?`,
|
||||
asString(req["name"]), typeVal, asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(asString(req["inIp"])), now, id)
|
||||
asString(req["name"]), typeVal, asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(inIp), now, id)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -2561,14 +2564,19 @@ func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRu
|
||||
if port <= 0 {
|
||||
return nil, errors.New("节点端口不能为空")
|
||||
}
|
||||
protocol := defaultString(target.Protocol, "tls")
|
||||
connector := map[string]interface{}{
|
||||
"type": "relay",
|
||||
}
|
||||
if isTLSTunnelProtocol(protocol) {
|
||||
connector["metadata"] = map[string]interface{}{"nodelay": true}
|
||||
}
|
||||
nodeItems = append(nodeItems, map[string]interface{}{
|
||||
"name": fmt.Sprintf("node_%d", idx+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, port)),
|
||||
"connector": map[string]interface{}{
|
||||
"type": "relay",
|
||||
},
|
||||
"name": fmt.Sprintf("node_%d", idx+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, port)),
|
||||
"connector": connector,
|
||||
"dialer": map[string]interface{}{
|
||||
"type": defaultString(target.Protocol, "tls"),
|
||||
"type": protocol,
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -2597,14 +2605,19 @@ func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode,
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
protocol := defaultString(chainNode.Protocol, "tls")
|
||||
handlerCfg := map[string]interface{}{
|
||||
"type": "relay",
|
||||
}
|
||||
if isTLSTunnelProtocol(protocol) {
|
||||
handlerCfg["metadata"] = map[string]interface{}{"nodelay": true}
|
||||
}
|
||||
service := map[string]interface{}{
|
||||
"name": fmt.Sprintf("%d_tls", tunnelID),
|
||||
"addr": fmt.Sprintf("%s:%d", node.TCPListenAddr, chainNode.Port),
|
||||
"handler": map[string]interface{}{
|
||||
"type": "relay",
|
||||
},
|
||||
"name": fmt.Sprintf("%d_tls", tunnelID),
|
||||
"addr": fmt.Sprintf("%s:%d", node.TCPListenAddr, chainNode.Port),
|
||||
"handler": handlerCfg,
|
||||
"listener": map[string]interface{}{
|
||||
"type": defaultString(chainNode.Protocol, "tls"),
|
||||
"type": protocol,
|
||||
},
|
||||
}
|
||||
if chainNode.ChainType == 2 {
|
||||
@@ -2650,6 +2663,10 @@ func nodeDisplayName(node *nodeRecord) string {
|
||||
return fmt.Sprintf("node_%d", node.ID)
|
||||
}
|
||||
|
||||
func isTLSTunnelProtocol(protocol string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(defaultString(protocol, "tls")), "tls")
|
||||
}
|
||||
|
||||
func nodeSupportsV4(node *nodeRecord) bool {
|
||||
if node == nil {
|
||||
return false
|
||||
@@ -2846,7 +2863,7 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac
|
||||
if nodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, NULL, NULL, 0, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '1', ?, NULL, NULL, 0, ?)`,
|
||||
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2865,7 +2882,7 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac
|
||||
return pickErr
|
||||
}
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, ?, ?, 0, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '3', ?, ?, ?, 0, ?)`,
|
||||
tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2886,7 +2903,7 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac
|
||||
return pickErr
|
||||
}
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, ?, ?, ?, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '2', ?, ?, ?, ?, ?)`,
|
||||
tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2972,7 +2989,7 @@ func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) {
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
|
||||
rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 ORDER BY inx ASC, id ASC`, tunnelID)
|
||||
rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = '1' ORDER BY inx ASC, id ASC`, tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
+234
-48
@@ -4,7 +4,7 @@ package store
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
@@ -98,7 +98,7 @@ func (db *DB) Begin() (*Tx, error) {
|
||||
func (db *DB) ExecReturningID(query string, args ...any) (int64, error) {
|
||||
q := db.rewrite(query)
|
||||
if db.dialect == DialectPostgres {
|
||||
q = strings.TrimRight(q, "; \t\n") + " RETURNING id"
|
||||
q = ensureReturningID(q)
|
||||
var id int64
|
||||
if err := db.raw.QueryRow(q, args...).Scan(&id); err != nil {
|
||||
return 0, err
|
||||
@@ -143,7 +143,7 @@ func (tx *Tx) Rollback() error { return tx.raw.Rollback() }
|
||||
func (tx *Tx) ExecReturningID(query string, args ...any) (int64, error) {
|
||||
q := rewriteQuery(tx.dialect, query)
|
||||
if tx.dialect == DialectPostgres {
|
||||
q = strings.TrimRight(q, "; \t\n") + " RETURNING id"
|
||||
q = ensureReturningID(q)
|
||||
var id int64
|
||||
if err := tx.raw.QueryRow(q, args...).Scan(&id); err != nil {
|
||||
return 0, err
|
||||
@@ -174,35 +174,15 @@ func rewriteQuery(dialect Dialect, query string) string {
|
||||
func rewriteUserIdentifier(query string) string {
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(query) + 16)
|
||||
inSingle := false
|
||||
inDouble := false
|
||||
i := 0
|
||||
for i < len(query) {
|
||||
ch := query[i]
|
||||
if ch == '\'' && !inDouble {
|
||||
if inSingle && i+1 < len(query) && query[i+1] == '\'' {
|
||||
buf.WriteByte(ch)
|
||||
buf.WriteByte(query[i+1])
|
||||
i += 2
|
||||
continue
|
||||
}
|
||||
inSingle = !inSingle
|
||||
buf.WriteByte(ch)
|
||||
i++
|
||||
continue
|
||||
}
|
||||
if ch == '"' && !inSingle {
|
||||
inDouble = !inDouble
|
||||
buf.WriteByte(ch)
|
||||
i++
|
||||
continue
|
||||
}
|
||||
if inSingle || inDouble {
|
||||
buf.WriteByte(ch)
|
||||
i++
|
||||
if end, ok := skipSQLProtectedSegment(query, i); ok {
|
||||
buf.WriteString(query[i:end])
|
||||
i = end
|
||||
continue
|
||||
}
|
||||
|
||||
ch := query[i]
|
||||
if isIdentifierChar(ch) {
|
||||
j := i + 1
|
||||
for j < len(query) && isIdentifierChar(query[j]) {
|
||||
@@ -238,39 +218,43 @@ func isIdentifierChar(ch byte) bool {
|
||||
}
|
||||
|
||||
func rewriteInsertOrIgnore(query string) string {
|
||||
upper := strings.ToUpper(query)
|
||||
idx := strings.Index(upper, "INSERT OR IGNORE INTO")
|
||||
if idx < 0 {
|
||||
start, end, ok := findKeywordSequenceOutside(query, []string{"INSERT", "OR", "IGNORE", "INTO"}, 0)
|
||||
if !ok {
|
||||
return query
|
||||
}
|
||||
prefix := query[:idx]
|
||||
suffix := query[idx+len("INSERT OR IGNORE INTO"):]
|
||||
result := prefix + "INSERT INTO" + suffix
|
||||
|
||||
trimmed := strings.TrimRight(result, "; \t\n")
|
||||
return trimmed + " ON CONFLICT DO NOTHING"
|
||||
rewritten := query[:start] + "INSERT INTO" + query[end:]
|
||||
rewritten = strings.TrimRight(rewritten, "; \t\n")
|
||||
|
||||
insertIntoEnd := start + len("INSERT INTO")
|
||||
if _, _, hasOnConflict := findKeywordSequenceOutside(rewritten, []string{"ON", "CONFLICT"}, insertIntoEnd); hasOnConflict {
|
||||
return rewritten
|
||||
}
|
||||
|
||||
if retStart, _, hasReturning := findKeywordSequenceOutside(rewritten, []string{"RETURNING"}, insertIntoEnd); hasReturning {
|
||||
prefix := strings.TrimRight(rewritten[:retStart], " \t\n")
|
||||
suffix := strings.TrimLeft(rewritten[retStart:], " \t\n")
|
||||
return prefix + " ON CONFLICT DO NOTHING " + suffix
|
||||
}
|
||||
|
||||
return rewritten + " ON CONFLICT DO NOTHING"
|
||||
}
|
||||
|
||||
func rewritePlaceholders(query string) string {
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(query) + 16)
|
||||
n := 1
|
||||
inString := false
|
||||
for i := 0; i < len(query); i++ {
|
||||
ch := query[i]
|
||||
if ch == '\'' {
|
||||
if inString && i+1 < len(query) && query[i+1] == '\'' {
|
||||
buf.WriteByte(ch)
|
||||
buf.WriteByte(query[i+1])
|
||||
i++
|
||||
continue
|
||||
}
|
||||
inString = !inString
|
||||
buf.WriteByte(ch)
|
||||
if end, ok := skipSQLProtectedSegment(query, i); ok {
|
||||
buf.WriteString(query[i:end])
|
||||
i = end - 1
|
||||
continue
|
||||
}
|
||||
if ch == '?' && !inString {
|
||||
buf.WriteString(fmt.Sprintf("$%d", n))
|
||||
|
||||
ch := query[i]
|
||||
if ch == '?' {
|
||||
buf.WriteByte('$')
|
||||
buf.WriteString(strconv.Itoa(n))
|
||||
n++
|
||||
continue
|
||||
}
|
||||
@@ -278,3 +262,205 @@ func rewritePlaceholders(query string) string {
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func ensureReturningID(query string) string {
|
||||
trimmed := strings.TrimRight(query, "; \t\n")
|
||||
if _, _, ok := findKeywordSequenceOutside(trimmed, []string{"RETURNING"}, 0); ok {
|
||||
return trimmed
|
||||
}
|
||||
return trimmed + " RETURNING id"
|
||||
}
|
||||
|
||||
func findKeywordSequenceOutside(query string, keywords []string, from int) (int, int, bool) {
|
||||
if len(keywords) == 0 {
|
||||
return 0, 0, false
|
||||
}
|
||||
if from < 0 {
|
||||
from = 0
|
||||
}
|
||||
if from >= len(query) {
|
||||
return 0, 0, false
|
||||
}
|
||||
|
||||
matched := 0
|
||||
seqStart := -1
|
||||
|
||||
for i := from; i < len(query); {
|
||||
if end, ok := skipSQLProtectedSegment(query, i); ok {
|
||||
i = end
|
||||
continue
|
||||
}
|
||||
|
||||
ch := query[i]
|
||||
if isIdentifierChar(ch) {
|
||||
j := i + 1
|
||||
for j < len(query) && isIdentifierChar(query[j]) {
|
||||
j++
|
||||
}
|
||||
tok := query[i:j]
|
||||
|
||||
if strings.EqualFold(tok, keywords[matched]) {
|
||||
if matched == 0 {
|
||||
seqStart = i
|
||||
}
|
||||
matched++
|
||||
if matched == len(keywords) {
|
||||
return seqStart, j, true
|
||||
}
|
||||
} else if strings.EqualFold(tok, keywords[0]) {
|
||||
seqStart = i
|
||||
matched = 1
|
||||
} else {
|
||||
matched = 0
|
||||
seqStart = -1
|
||||
}
|
||||
|
||||
i = j
|
||||
continue
|
||||
}
|
||||
|
||||
if !isSQLSpace(ch) {
|
||||
matched = 0
|
||||
seqStart = -1
|
||||
}
|
||||
i++
|
||||
}
|
||||
|
||||
return 0, 0, false
|
||||
}
|
||||
|
||||
func skipSQLProtectedSegment(query string, i int) (int, bool) {
|
||||
if i < 0 || i >= len(query) {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
switch query[i] {
|
||||
case '\'':
|
||||
return skipSingleQuotedLiteral(query, i), true
|
||||
case '"':
|
||||
return skipDoubleQuotedIdentifier(query, i), true
|
||||
case '-':
|
||||
if i+1 < len(query) && query[i+1] == '-' {
|
||||
return skipLineComment(query, i), true
|
||||
}
|
||||
case '/':
|
||||
if i+1 < len(query) && query[i+1] == '*' {
|
||||
return skipBlockComment(query, i), true
|
||||
}
|
||||
case '$':
|
||||
if end, ok := skipDollarQuotedLiteral(query, i); ok {
|
||||
return end, true
|
||||
}
|
||||
}
|
||||
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func skipSingleQuotedLiteral(query string, i int) int {
|
||||
for j := i + 1; j < len(query); j++ {
|
||||
if query[j] != '\'' {
|
||||
continue
|
||||
}
|
||||
if j+1 < len(query) && query[j+1] == '\'' {
|
||||
j++
|
||||
continue
|
||||
}
|
||||
return j + 1
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipDoubleQuotedIdentifier(query string, i int) int {
|
||||
for j := i + 1; j < len(query); j++ {
|
||||
if query[j] != '"' {
|
||||
continue
|
||||
}
|
||||
if j+1 < len(query) && query[j+1] == '"' {
|
||||
j++
|
||||
continue
|
||||
}
|
||||
return j + 1
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipLineComment(query string, i int) int {
|
||||
for j := i + 2; j < len(query); j++ {
|
||||
if query[j] == '\n' {
|
||||
return j
|
||||
}
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipBlockComment(query string, i int) int {
|
||||
depth := 1
|
||||
for j := i + 2; j < len(query)-1; j++ {
|
||||
if query[j] == '/' && query[j+1] == '*' {
|
||||
depth++
|
||||
j++
|
||||
continue
|
||||
}
|
||||
if query[j] == '*' && query[j+1] == '/' {
|
||||
depth--
|
||||
j++
|
||||
if depth == 0 {
|
||||
return j + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipDollarQuotedLiteral(query string, i int) (int, bool) {
|
||||
if i < 0 || i >= len(query) || query[i] != '$' {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
if i+1 >= len(query) {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
var endTag int
|
||||
if query[i+1] == '$' {
|
||||
endTag = i + 1
|
||||
} else {
|
||||
if !isDollarTagStart(query[i+1]) {
|
||||
return 0, false
|
||||
}
|
||||
j := i + 2
|
||||
for j < len(query) && isDollarTagChar(query[j]) {
|
||||
j++
|
||||
}
|
||||
if j >= len(query) || query[j] != '$' {
|
||||
return 0, false
|
||||
}
|
||||
endTag = j
|
||||
}
|
||||
|
||||
tag := query[i : endTag+1]
|
||||
if closeIdx := strings.Index(query[endTag+1:], tag); closeIdx >= 0 {
|
||||
return endTag + 1 + closeIdx + len(tag), true
|
||||
}
|
||||
return len(query), true
|
||||
}
|
||||
|
||||
func isDollarTagStart(ch byte) bool {
|
||||
return ch == '_' || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z')
|
||||
}
|
||||
|
||||
func isDollarTagChar(ch byte) bool {
|
||||
if isDollarTagStart(ch) {
|
||||
return true
|
||||
}
|
||||
return ch >= '0' && ch <= '9'
|
||||
}
|
||||
|
||||
func isSQLSpace(ch byte) bool {
|
||||
switch ch {
|
||||
case ' ', '\t', '\n', '\r', '\f':
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
package store
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestRewritePlaceholdersSkipsProtectedSegments(t *testing.T) {
|
||||
q := `SELECT ?, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ?
|
||||
FROM t /* block ? */ WHERE id = ?`
|
||||
got := rewritePlaceholders(q)
|
||||
want := `SELECT $1, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ?
|
||||
FROM t /* block ? */ WHERE id = $2`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreBasic(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreBeforeReturning(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO x(a) VALUES(?) RETURNING id`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `INSERT INTO x(a) VALUES(?) ON CONFLICT DO NOTHING RETURNING id`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreNotDuplicatingOnConflict(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `INSERT INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureReturningID(t *testing.T) {
|
||||
if got := ensureReturningID(`INSERT INTO x(a) VALUES($1)`); got != `INSERT INTO x(a) VALUES($1) RETURNING id` {
|
||||
t.Fatalf("missing RETURNING append: %s", got)
|
||||
}
|
||||
if got := ensureReturningID(`INSERT INTO x(a) VALUES($1) RETURNING other_id`); got != `INSERT INTO x(a) VALUES($1) RETURNING other_id` {
|
||||
t.Fatalf("RETURNING should not be duplicated: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteUserIdentifierSafety(t *testing.T) {
|
||||
q := `SELECT user, user_id, 'user', "user", note FROM user -- user
|
||||
WHERE owner='user'`
|
||||
got := rewriteUserIdentifier(q)
|
||||
want := `SELECT "user", user_id, 'user', "user", note FROM "user" -- user
|
||||
WHERE owner='user'`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteQueryPostgresPipeline(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO user(name, note) VALUES(?, '?')`
|
||||
got := rewriteQuery(DialectPostgres, q)
|
||||
want := `INSERT INTO "user"(name, note) VALUES($1, '?') ON CONFLICT DO NOTHING`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreSkipsStringLiteral(t *testing.T) {
|
||||
q := `SELECT 'INSERT OR IGNORE INTO t(a) VALUES(?)' AS q`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
if got != q {
|
||||
t.Fatalf("string literal should stay unchanged\nwant: %s\ngot: %s", q, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreSkipsCommentedKeyword(t *testing.T) {
|
||||
q := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?)
|
||||
INSERT OR IGNORE INTO real_t(a) VALUES(?)`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?)
|
||||
INSERT INTO real_t(a) VALUES(?) ON CONFLICT DO NOTHING`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewritePlaceholdersSkipsNestedBlockComment(t *testing.T) {
|
||||
q := `SELECT ? /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = ?`
|
||||
got := rewritePlaceholders(q)
|
||||
want := `SELECT $1 /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = $2`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewritePlaceholdersSkipsUnterminatedBlockComment(t *testing.T) {
|
||||
q := `SELECT ? /* unterminated ? comment`
|
||||
got := rewritePlaceholders(q)
|
||||
want := `SELECT $1 /* unterminated ? comment`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteUserIdentifierSkipsDollarQuotedAndComment(t *testing.T) {
|
||||
q := `SELECT user, $$user ?$$ AS body, col FROM user /* user */ -- user`
|
||||
got := rewriteUserIdentifier(q)
|
||||
want := `SELECT "user", $$user ?$$ AS body, col FROM "user" /* user */ -- user`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
@@ -400,7 +400,7 @@ func (r *Repository) GetUserPackageForwards(userID int64) ([]UserForwardDetail,
|
||||
}
|
||||
|
||||
rows, err := r.db.Query(`
|
||||
SELECT f.id, f.name, f.tunnel_id, t.name, f.remote_addr, f.in_flow, f.out_flow, f.status, f.created_time
|
||||
SELECT f.id, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, f.in_flow, f.out_flow, f.status, f.created_time
|
||||
FROM forward f
|
||||
LEFT JOIN tunnel t ON t.id = f.tunnel_id
|
||||
WHERE f.user_id = ?
|
||||
@@ -719,7 +719,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
}
|
||||
|
||||
rows, err := r.db.Query(`
|
||||
SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, t.name, f.remote_addr, f.strategy,
|
||||
SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, f.strategy,
|
||||
f.in_flow, f.out_flow, f.created_time, f.status, f.inx
|
||||
FROM forward f
|
||||
LEFT JOIN tunnel t ON t.id = f.tunnel_id
|
||||
@@ -897,9 +897,9 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
}
|
||||
|
||||
chainRows, err := r.db.Query(`
|
||||
SELECT tunnel_id, chain_type, node_id, protocol, strategy, COALESCE(inx, 0)
|
||||
SELECT tunnel_id, CAST(chain_type AS INTEGER), node_id, protocol, strategy, COALESCE(inx, 0)
|
||||
FROM chain_tunnel
|
||||
ORDER BY tunnel_id ASC, chain_type ASC, inx ASC, id ASC
|
||||
ORDER BY tunnel_id ASC, CAST(chain_type AS INTEGER) ASC, inx ASC, id ASC
|
||||
`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1297,11 +1297,32 @@ func bootstrapSchema(db *store.DB, schemaSQL, seedSQL string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
const currentSchemaVersion = 1
|
||||
|
||||
func getSchemaVersion(db *store.DB) int {
|
||||
_, _ = db.Exec(`CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL DEFAULT 0)`)
|
||||
var v int
|
||||
if err := db.QueryRow(`SELECT version FROM schema_version LIMIT 1`).Scan(&v); err != nil {
|
||||
_, _ = db.Exec(`INSERT INTO schema_version(version) VALUES(0)`)
|
||||
return 0
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func setSchemaVersion(db *store.DB, v int) {
|
||||
_, _ = db.Exec(`UPDATE schema_version SET version = ?`, v)
|
||||
}
|
||||
|
||||
func migrateSchema(db *store.DB) error {
|
||||
if db == nil {
|
||||
return errors.New("nil db")
|
||||
}
|
||||
|
||||
ver := getSchemaVersion(db)
|
||||
if ver >= currentSchemaVersion {
|
||||
return nil
|
||||
}
|
||||
|
||||
ensureColumn := func(table, col, typ string) {
|
||||
var dummy interface{}
|
||||
err := db.QueryRow(fmt.Sprintf("SELECT %s FROM %s LIMIT 1", col, table)).Scan(&dummy)
|
||||
@@ -1351,6 +1372,7 @@ func migrateSchema(db *store.DB) error {
|
||||
return err
|
||||
}
|
||||
}
|
||||
setSchemaVersion(db, currentSchemaVersion)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -54,10 +54,10 @@ var needWrap = false
|
||||
|
||||
// SetProtocolBlock sets protocol blocking switches and recomputes wrapper need
|
||||
func SetProtocolBlock(httpOn int, tlsOn int, socksOn int) {
|
||||
isHttp = httpOn
|
||||
isTls = tlsOn
|
||||
isSocks = socksOn
|
||||
needWrap = isTls+isSocks+isHttp > 0
|
||||
isHttp = httpOn
|
||||
isTls = tlsOn
|
||||
isSocks = socksOn
|
||||
needWrap = isTls+isSocks+isHttp > 0
|
||||
}
|
||||
|
||||
type Option func(opts *options)
|
||||
@@ -292,7 +292,9 @@ func (s *defaultService) Serve() error {
|
||||
}
|
||||
|
||||
if err := s.handler.Handle(ctx, conn); err != nil {
|
||||
log.Error(err)
|
||||
if !errors.Is(err, net.ErrClosed) {
|
||||
log.Error(err)
|
||||
}
|
||||
if v := xmetrics.GetCounter(xmetrics.MetricServiceHandlerErrorsCounter,
|
||||
metrics.Labels{"service": s.name, "client": clientIP}); v != nil {
|
||||
v.Inc()
|
||||
@@ -403,12 +405,12 @@ func (s *defaultService) observeStats(ctx context.Context) {
|
||||
TotalErrs: st.Get(stats.KindTotalErrs),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
// 将流量累积到全局管理器,而不是立即上报
|
||||
if outputBytes > 0 || inputBytes > 0 {
|
||||
globalManager := GetGlobalTrafficManager()
|
||||
globalManager.AddTraffic(s.name, int64(outputBytes), int64(inputBytes))
|
||||
|
||||
|
||||
// 立即重置流量计数(因为已经记录到全局管理器中)
|
||||
if xstats, ok := st.(*xstats.Stats); ok {
|
||||
xstats.ResetTraffic(st.Get(stats.KindInputBytes)-inputBytes, st.Get(stats.KindOutputBytes)-outputBytes)
|
||||
|
||||
Reference in New Issue
Block a user