import.go 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374
  1. package sqliteimport
  2. import (
  3. "database/sql"
  4. "encoding/hex"
  5. "fmt"
  6. "os"
  7. "regexp"
  8. "strings"
  9. "github.com/danfragoso/pizzasql-next/pkg/executor"
  10. "github.com/danfragoso/pizzasql-next/pkg/lexer"
  11. "github.com/danfragoso/pizzasql-next/pkg/parser"
  12. _ "modernc.org/sqlite"
  13. )
  14. const insertBatchSize = 500
  15. // ImportOptions configures SQLite import behavior.
  16. type ImportOptions struct {
  17. CreateTables bool // Create tables from the source schema (default true)
  18. IgnoreErrors bool // Continue on individual row/statement errors
  19. TableFilter []string // If non-empty, only import these tables
  20. }
  21. // DefaultImportOptions returns sensible defaults.
  22. func DefaultImportOptions() ImportOptions {
  23. return ImportOptions{
  24. CreateTables: true,
  25. IgnoreErrors: false,
  26. }
  27. }
  28. // ImportResult contains the results of an import operation.
  29. type ImportResult struct {
  30. TablesCreated []string `json:"tablesCreated"`
  31. TablesImported []string `json:"tablesImported"`
  32. RowsInserted int64 `json:"rowsInserted"`
  33. IndexesCreated int `json:"indexesCreated"`
  34. Errors []string `json:"errors,omitempty"`
  35. }
  36. // ImportSQLiteFile imports a SQLite .db file into a PizzaSQL executor.
  37. func ImportSQLiteFile(path string, exec *executor.Executor, opts ImportOptions) (*ImportResult, error) {
  38. db, err := sql.Open("sqlite", path+"?mode=ro")
  39. if err != nil {
  40. return nil, fmt.Errorf("open sqlite file: %w", err)
  41. }
  42. defer db.Close()
  43. if err := db.Ping(); err != nil {
  44. return nil, fmt.Errorf("cannot read sqlite file: %w", err)
  45. }
  46. return importFromDB(db, exec, opts)
  47. }
  48. // ImportSQLiteBytes imports a SQLite database from raw bytes (e.g. from an HTTP upload).
  49. // It writes to a temporary file, imports, then removes the file.
  50. func ImportSQLiteBytes(data []byte, exec *executor.Executor, opts ImportOptions) (*ImportResult, error) {
  51. tmp, err := os.CreateTemp("", "pizzasql-sqlite-*.db")
  52. if err != nil {
  53. return nil, fmt.Errorf("create temp file: %w", err)
  54. }
  55. tmpPath := tmp.Name()
  56. defer os.Remove(tmpPath)
  57. if _, err := tmp.Write(data); err != nil {
  58. tmp.Close()
  59. return nil, fmt.Errorf("write temp file: %w", err)
  60. }
  61. tmp.Close()
  62. return ImportSQLiteFile(tmpPath, exec, opts)
  63. }
  64. func importFromDB(db *sql.DB, exec *executor.Executor, opts ImportOptions) (*ImportResult, error) {
  65. result := &ImportResult{
  66. TablesCreated: []string{},
  67. TablesImported: []string{},
  68. Errors: []string{},
  69. }
  70. // Load table list and DDL from sqlite_master
  71. type tableEntry struct {
  72. name string
  73. ddl string
  74. }
  75. var tables []tableEntry
  76. rows, err := db.Query(`SELECT name, sql FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY rowid`)
  77. if err != nil {
  78. return nil, fmt.Errorf("query sqlite_master: %w", err)
  79. }
  80. defer rows.Close()
  81. for rows.Next() {
  82. var name string
  83. var ddlNull sql.NullString
  84. if err := rows.Scan(&name, &ddlNull); err != nil {
  85. continue
  86. }
  87. if !ddlNull.Valid || ddlNull.String == "" {
  88. continue
  89. }
  90. tables = append(tables, tableEntry{name: name, ddl: ddlNull.String})
  91. }
  92. rows.Close()
  93. // Apply table filter
  94. if len(opts.TableFilter) > 0 {
  95. filter := make(map[string]bool, len(opts.TableFilter))
  96. for _, t := range opts.TableFilter {
  97. filter[strings.ToLower(t)] = true
  98. }
  99. filtered := tables[:0]
  100. for _, t := range tables {
  101. if filter[strings.ToLower(t.name)] {
  102. filtered = append(filtered, t)
  103. }
  104. }
  105. tables = filtered
  106. }
  107. // Create tables
  108. if opts.CreateTables {
  109. for _, t := range tables {
  110. ddl := sanitizeDDL(t.ddl)
  111. if err := execStatement(exec, ddl); err != nil {
  112. msg := fmt.Sprintf("create table %s: %v", t.name, err)
  113. result.Errors = append(result.Errors, msg)
  114. if !opts.IgnoreErrors {
  115. return result, fmt.Errorf("%s", msg)
  116. }
  117. continue
  118. }
  119. result.TablesCreated = append(result.TablesCreated, t.name)
  120. }
  121. }
  122. // Import indexes
  123. idxRows, err := db.Query(`SELECT sql FROM sqlite_master WHERE type='index' AND sql IS NOT NULL AND name NOT LIKE 'sqlite_%'`)
  124. if err == nil {
  125. defer idxRows.Close()
  126. for idxRows.Next() {
  127. var idxSQL string
  128. if err := idxRows.Scan(&idxSQL); err != nil {
  129. continue
  130. }
  131. idxSQL = sanitizeDDL(idxSQL)
  132. if err := execStatement(exec, idxSQL); err != nil {
  133. result.Errors = append(result.Errors, fmt.Sprintf("create index: %v", err))
  134. } else {
  135. result.IndexesCreated++
  136. }
  137. }
  138. idxRows.Close()
  139. }
  140. // Import rows per table
  141. for _, t := range tables {
  142. n, err := importTableRows(db, exec, t.name, opts.IgnoreErrors)
  143. if err != nil {
  144. msg := fmt.Sprintf("import rows for %s: %v", t.name, err)
  145. result.Errors = append(result.Errors, msg)
  146. if !opts.IgnoreErrors {
  147. return result, fmt.Errorf("%s", msg)
  148. }
  149. continue
  150. }
  151. result.TablesImported = append(result.TablesImported, t.name)
  152. result.RowsInserted += n
  153. }
  154. return result, nil
  155. }
  156. func importTableRows(db *sql.DB, exec *executor.Executor, table string, ignoreErrors bool) (int64, error) {
  157. rows, err := db.Query(fmt.Sprintf(`SELECT * FROM %q`, table))
  158. if err != nil {
  159. return 0, err
  160. }
  161. defer rows.Close()
  162. cols, err := rows.Columns()
  163. if err != nil {
  164. return 0, err
  165. }
  166. if len(cols) == 0 {
  167. return 0, nil
  168. }
  169. var total int64
  170. var batch []string
  171. flush := func() error {
  172. if len(batch) == 0 {
  173. return nil
  174. }
  175. // Build multi-row INSERT
  176. colList := quoteIdentList(cols)
  177. sql := fmt.Sprintf("INSERT INTO %s (%s) VALUES %s",
  178. quoteIdent(table), colList, strings.Join(batch, ", "))
  179. if err := execStatement(exec, sql); err != nil {
  180. return err
  181. }
  182. total += int64(len(batch))
  183. batch = batch[:0]
  184. return nil
  185. }
  186. vals := make([]interface{}, len(cols))
  187. ptrs := make([]interface{}, len(cols))
  188. for i := range vals {
  189. ptrs[i] = &vals[i]
  190. }
  191. for rows.Next() {
  192. if err := rows.Scan(ptrs...); err != nil {
  193. if ignoreErrors {
  194. continue
  195. }
  196. return total, err
  197. }
  198. batch = append(batch, rowToValueList(vals))
  199. if len(batch) >= insertBatchSize {
  200. if err := flush(); err != nil {
  201. if ignoreErrors {
  202. batch = batch[:0]
  203. continue
  204. }
  205. return total, err
  206. }
  207. }
  208. }
  209. if err := rows.Err(); err != nil {
  210. return total, err
  211. }
  212. if err := flush(); err != nil {
  213. return total, err
  214. }
  215. return total, nil
  216. }
  217. // rowToValueList converts a row of Go values into a SQL VALUES tuple string.
  218. func rowToValueList(vals []interface{}) string {
  219. parts := make([]string, len(vals))
  220. for i, v := range vals {
  221. parts[i] = sqlLiteral(v)
  222. }
  223. return "(" + strings.Join(parts, ", ") + ")"
  224. }
  225. // sqlLiteral converts a Go value (from the sqlite driver) to a SQL literal string.
  226. func sqlLiteral(v interface{}) string {
  227. if v == nil {
  228. return "NULL"
  229. }
  230. switch val := v.(type) {
  231. case int64:
  232. return fmt.Sprintf("%d", val)
  233. case float64:
  234. return fmt.Sprintf("%g", val)
  235. case string:
  236. return "'" + strings.ReplaceAll(val, "'", "''") + "'"
  237. case []byte:
  238. // Store blobs as hex text strings (PizzaSQL has no X'' literal support)
  239. return "'" + hex.EncodeToString(val) + "'"
  240. case bool:
  241. if val {
  242. return "1"
  243. }
  244. return "0"
  245. default:
  246. s := fmt.Sprintf("%v", val)
  247. return "'" + strings.ReplaceAll(s, "'", "''") + "'"
  248. }
  249. }
  250. func quoteIdent(s string) string {
  251. return `"` + strings.ReplaceAll(s, `"`, `""`) + `"`
  252. }
  253. func quoteIdentList(cols []string) string {
  254. parts := make([]string, len(cols))
  255. for i, c := range cols {
  256. parts[i] = quoteIdent(c)
  257. }
  258. return strings.Join(parts, ", ")
  259. }
  260. func execStatement(exec *executor.Executor, sql string) error {
  261. l := lexer.New(sql)
  262. p := parser.New(l)
  263. stmt, err := p.Parse()
  264. if err != nil {
  265. return fmt.Errorf("parse: %w", err)
  266. }
  267. _, err = exec.Execute(stmt)
  268. return err
  269. }
  270. // pizzasqlReservedKeywords is the set of tokens that PizzaSQL reserves but are
  271. // commonly used as column names in SQLite schemas.
  272. var pizzasqlReservedKeywords = map[string]bool{
  273. "key": true, "value": true, "type": true, "name": true,
  274. "index": true, "view": true, "table": true, "column": true,
  275. "group": true, "order": true, "range": true, "match": true,
  276. }
  277. var (
  278. reAutoincrement = regexp.MustCompile(`(?i)\bAUTOINCREMENT\b`)
  279. reWithoutRowid = regexp.MustCompile(`(?i)\bWITHOUT\s+ROWID\b`)
  280. reStrict = regexp.MustCompile(`(?i),?\s*\bSTRICT\b`)
  281. // REFERENCES x(y) ON DELETE/UPDATE action — strip whole inline FK clause.
  282. // Use \w+ (not \S+) so the trailing comma of the column is preserved.
  283. reInlineRefs = regexp.MustCompile(`(?i)\bREFERENCES\s+\w+\s*(?:\([^)]*\))?\s*(?:(?:ON\s+(?:DELETE|UPDATE)\s+(?:CASCADE|SET\s+NULL|SET\s+DEFAULT|RESTRICT|NO\s+ACTION))\s*)*`)
  284. // Table-level FOREIGN KEY constraint lines
  285. reTableFK = regexp.MustCompile(`(?i),?\s*FOREIGN\s+KEY\s*\([^)]*\)\s*REFERENCES\s+\w+\s*(?:\([^)]*\))?\s*(?:(?:ON\s+(?:DELETE|UPDATE)\s+(?:CASCADE|SET\s+NULL|SET\s+DEFAULT|RESTRICT|NO\s+ACTION))\s*)*`)
  286. // Table-level CHECK constraints
  287. reTableCheck = regexp.MustCompile(`(?i),?\s*CHECK\s*\([^)]*\)`)
  288. reOnConflict = regexp.MustCompile(`(?i)\bON\s+CONFLICT\s+\w+`)
  289. // Complex DEFAULT expressions: DEFAULT (...) — strip entirely, keep no default
  290. reComplexDefault = regexp.MustCompile(`(?i)\bDEFAULT\s*\([^)]*\)`)
  291. // Trailing comma before closing paren
  292. reTableTrailing = regexp.MustCompile(`(?m),\s*\)`)
  293. // DESC/ASC in index column lists
  294. reIndexColOrder = regexp.MustCompile(`(?i)\b(ASC|DESC)\b`)
  295. // Column name (first word) followed by a type keyword on each column line
  296. reColumnName = regexp.MustCompile(`(?m)^\s{1,}(\w+)(\s+)`)
  297. )
  298. // sanitizeDDL strips SQLite-specific clauses that PizzaSQL doesn't support.
  299. func sanitizeDDL(ddl string) string {
  300. ddl = reAutoincrement.ReplaceAllString(ddl, "")
  301. ddl = reWithoutRowid.ReplaceAllString(ddl, "")
  302. ddl = reStrict.ReplaceAllString(ddl, "")
  303. ddl = reTableFK.ReplaceAllString(ddl, "")
  304. ddl = reTableCheck.ReplaceAllString(ddl, "")
  305. ddl = reComplexDefault.ReplaceAllString(ddl, "")
  306. ddl = reInlineRefs.ReplaceAllString(ddl, "")
  307. ddl = reOnConflict.ReplaceAllString(ddl, "")
  308. // Strip ASC/DESC from index column lists
  309. upper := strings.ToUpper(ddl)
  310. if strings.Contains(upper, "CREATE INDEX") || strings.Contains(upper, "CREATE UNIQUE INDEX") {
  311. ddl = reIndexColOrder.ReplaceAllString(ddl, "")
  312. }
  313. // Quote column names that clash with PizzaSQL reserved keywords
  314. ddl = reColumnName.ReplaceAllStringFunc(ddl, func(m string) string {
  315. // Extract leading whitespace, word, trailing whitespace
  316. sub := reColumnName.FindStringSubmatch(m)
  317. if len(sub) < 3 {
  318. return m
  319. }
  320. word, ws := sub[1], sub[2]
  321. if pizzasqlReservedKeywords[strings.ToLower(word)] {
  322. leading := m[:len(m)-len(word)-len(ws)]
  323. return leading + `"` + word + `"` + ws
  324. }
  325. return m
  326. })
  327. // Clean up trailing commas before closing paren
  328. ddl = reTableTrailing.ReplaceAllStringFunc(ddl, func(s string) string {
  329. return ")"
  330. })
  331. return strings.TrimSpace(ddl)
  332. }