export.go 6.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283
  1. package sqlexport
  2. import (
  3. "encoding/hex"
  4. "fmt"
  5. "sort"
  6. "strings"
  7. "time"
  8. "github.com/danfragoso/pizzasql-next/pkg/storage"
  9. )
  10. // ExportOptions configures export behavior.
  11. type ExportOptions struct {
  12. Tables []string // Specific tables to export (empty = all tables)
  13. IncludeData bool // Include INSERT statements (default: true)
  14. DropTables bool // Add DROP TABLE IF EXISTS before CREATE
  15. }
  16. // DefaultExportOptions returns sensible defaults.
  17. func DefaultExportOptions() ExportOptions {
  18. return ExportOptions{
  19. Tables: nil,
  20. IncludeData: true,
  21. DropTables: false,
  22. }
  23. }
  24. // ExportDatabase exports an entire database to SQL text.
  25. func ExportDatabase(schema *storage.SchemaManager, table *storage.TableManager, opts ExportOptions) (string, error) {
  26. var sb strings.Builder
  27. // Write header
  28. sb.WriteString("-- PizzaSQL Export\n")
  29. sb.WriteString(fmt.Sprintf("-- Database: %s\n", schema.GetDatabaseName()))
  30. sb.WriteString(fmt.Sprintf("-- Date: %s\n", time.Now().UTC().Format(time.RFC3339)))
  31. sb.WriteString("\n")
  32. // Get tables to export
  33. tables := opts.Tables
  34. if len(tables) == 0 {
  35. var err error
  36. tables, err = schema.ListTables()
  37. if err != nil {
  38. return "", fmt.Errorf("failed to list tables: %w", err)
  39. }
  40. }
  41. // Sort tables for consistent output
  42. sort.Strings(tables)
  43. // Export each table
  44. for i, tableName := range tables {
  45. tableSQL, err := ExportTable(schema, table, tableName, opts)
  46. if err != nil {
  47. return "", fmt.Errorf("failed to export table %s: %w", tableName, err)
  48. }
  49. sb.WriteString(tableSQL)
  50. // Add separator between tables
  51. if i < len(tables)-1 {
  52. sb.WriteString("\n")
  53. }
  54. }
  55. return sb.String(), nil
  56. }
  57. // ExportTable exports a single table to SQL text.
  58. func ExportTable(schema *storage.SchemaManager, table *storage.TableManager, tableName string, opts ExportOptions) (string, error) {
  59. var sb strings.Builder
  60. // Get table schema
  61. tableSchema, err := schema.GetSchema(tableName)
  62. if err != nil {
  63. return "", fmt.Errorf("failed to get schema for table %s: %w", tableName, err)
  64. }
  65. // Write DROP TABLE if requested
  66. if opts.DropTables {
  67. sb.WriteString(fmt.Sprintf("DROP TABLE IF EXISTS %s;\n", quoteIdentifier(tableName)))
  68. }
  69. // Write CREATE TABLE
  70. createSQL := generateCreateTable(tableSchema)
  71. sb.WriteString(createSQL)
  72. sb.WriteString("\n")
  73. // Write INSERT statements if requested
  74. if opts.IncludeData {
  75. rows, err := table.Select(tableName, nil)
  76. if err != nil {
  77. return "", fmt.Errorf("failed to select data from table %s: %w", tableName, err)
  78. }
  79. if len(rows) > 0 {
  80. sb.WriteString("\n")
  81. insertSQL := generateInserts(tableName, tableSchema, rows)
  82. sb.WriteString(insertSQL)
  83. }
  84. }
  85. return sb.String(), nil
  86. }
  87. // generateCreateTable generates a CREATE TABLE statement from schema.
  88. func generateCreateTable(schema *storage.Schema) string {
  89. var sb strings.Builder
  90. sb.WriteString(fmt.Sprintf("CREATE TABLE %s (\n", quoteIdentifier(schema.Name)))
  91. // Generate column definitions
  92. for i, col := range schema.Columns {
  93. // Skip internal _rowid_ column
  94. if col.Name == "_rowid_" {
  95. continue
  96. }
  97. sb.WriteString(" ")
  98. sb.WriteString(quoteIdentifier(col.Name))
  99. sb.WriteString(" ")
  100. sb.WriteString(col.Type)
  101. // Add PRIMARY KEY constraint
  102. if col.PrimaryKey {
  103. sb.WriteString(" PRIMARY KEY")
  104. }
  105. // Add NOT NULL constraint
  106. if !col.Nullable && !col.PrimaryKey {
  107. sb.WriteString(" NOT NULL")
  108. }
  109. // Add DEFAULT value
  110. if col.Default != nil {
  111. sb.WriteString(" DEFAULT ")
  112. sb.WriteString(formatValue(col.Default, col.Type))
  113. }
  114. // Add comma if not last column
  115. if i < len(schema.Columns)-1 {
  116. // Check if next column is _rowid_
  117. if i+1 < len(schema.Columns) && schema.Columns[i+1].Name != "_rowid_" {
  118. sb.WriteString(",")
  119. } else if i+2 < len(schema.Columns) {
  120. sb.WriteString(",")
  121. }
  122. }
  123. sb.WriteString("\n")
  124. }
  125. sb.WriteString(");")
  126. return sb.String()
  127. }
  128. // generateInserts generates INSERT statements for rows.
  129. func generateInserts(tableName string, schema *storage.Schema, rows []storage.Row) string {
  130. var sb strings.Builder
  131. // Get column names (excluding _rowid_)
  132. var columns []string
  133. var colTypes []string
  134. for _, col := range schema.Columns {
  135. if col.Name != "_rowid_" {
  136. columns = append(columns, col.Name)
  137. colTypes = append(colTypes, col.Type)
  138. }
  139. }
  140. // Generate INSERT for each row
  141. for _, row := range rows {
  142. sb.WriteString(fmt.Sprintf("INSERT INTO %s (", quoteIdentifier(tableName)))
  143. // Column names
  144. for i, col := range columns {
  145. if i > 0 {
  146. sb.WriteString(", ")
  147. }
  148. sb.WriteString(quoteIdentifier(col))
  149. }
  150. sb.WriteString(") VALUES (")
  151. // Values
  152. for i, col := range columns {
  153. if i > 0 {
  154. sb.WriteString(", ")
  155. }
  156. value := row[col]
  157. sb.WriteString(formatValue(value, colTypes[i]))
  158. }
  159. sb.WriteString(");\n")
  160. }
  161. return sb.String()
  162. }
  163. // formatValue formats a Go value as a SQL literal.
  164. func formatValue(value interface{}, colType string) string {
  165. if value == nil {
  166. return "NULL"
  167. }
  168. switch v := value.(type) {
  169. case string:
  170. return formatString(v)
  171. case float64:
  172. // Check if it's actually an integer
  173. if strings.Contains(strings.ToUpper(colType), "INT") {
  174. return fmt.Sprintf("%d", int64(v))
  175. }
  176. return fmt.Sprintf("%g", v)
  177. case int64:
  178. return fmt.Sprintf("%d", v)
  179. case int:
  180. return fmt.Sprintf("%d", v)
  181. case bool:
  182. if v {
  183. return "1"
  184. }
  185. return "0"
  186. case []byte:
  187. return fmt.Sprintf("X'%s'", hex.EncodeToString(v))
  188. default:
  189. // Fallback: treat as string
  190. return formatString(fmt.Sprintf("%v", v))
  191. }
  192. }
  193. // formatString formats a string as a SQL string literal with proper escaping.
  194. func formatString(s string) string {
  195. // Escape single quotes by doubling them
  196. escaped := strings.ReplaceAll(s, "'", "''")
  197. return fmt.Sprintf("'%s'", escaped)
  198. }
  199. // quoteIdentifier quotes a SQL identifier if needed.
  200. func quoteIdentifier(name string) string {
  201. // Check if identifier needs quoting
  202. needsQuote := false
  203. // Check for reserved words or special characters
  204. lower := strings.ToLower(name)
  205. reserved := map[string]bool{
  206. "table": true, "select": true, "insert": true, "update": true,
  207. "delete": true, "create": true, "drop": true, "index": true,
  208. "from": true, "where": true, "and": true, "or": true,
  209. "order": true, "by": true, "group": true, "having": true,
  210. "limit": true, "offset": true, "join": true, "on": true,
  211. "as": true, "null": true, "not": true, "in": true,
  212. "like": true, "between": true, "is": true, "primary": true,
  213. "key": true, "unique": true, "default": true, "values": true,
  214. }
  215. if reserved[lower] {
  216. needsQuote = true
  217. }
  218. // Check for special characters
  219. for _, c := range name {
  220. if !((c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
  221. (c >= '0' && c <= '9') || c == '_') {
  222. needsQuote = true
  223. break
  224. }
  225. }
  226. // Check if starts with digit
  227. if len(name) > 0 && name[0] >= '0' && name[0] <= '9' {
  228. needsQuote = true
  229. }
  230. if needsQuote {
  231. // Use double quotes and escape any existing double quotes
  232. escaped := strings.ReplaceAll(name, "\"", "\"\"")
  233. return fmt.Sprintf("\"%s\"", escaped)
  234. }
  235. return name
  236. }