Skip to content

Commit 524749e

Browse files
authored
Merge pull request #127 from planetscale/add-null-first
Handle boolean and null time values
2 parents 9d02f43 + 2a2c484 commit 524749e

3 files changed

Lines changed: 149 additions & 21 deletions

File tree

cmd/internal/planetscale_edge_database.go

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -421,7 +421,7 @@ func (p PlanetScaleEdgeDatabase) sync(ctx context.Context, syncMode string, tc *
421421
}
422422
sqlResult.Rows = append(sqlResult.Rows, row)
423423
// Results queued to Airbyte here, and flushed at the end of sync()
424-
p.printQueryResult(sqlResult, keyspaceOrDatabase, s.Name)
424+
p.printQueryResult(sqlResult, keyspaceOrDatabase, s.Name, &ps)
425425
}
426426
}
427427
}
@@ -509,8 +509,8 @@ func (p PlanetScaleEdgeDatabase) initializeVTGateClient(ctx context.Context, ps
509509

510510
// printQueryResult will pretty-print an AirbyteRecordMessage to the logger.
511511
// Copied from vtctl/query.go
512-
func (p PlanetScaleEdgeDatabase) printQueryResult(qr *sqltypes.Result, tableNamespace, tableName string) {
513-
data := QueryResultToRecords(qr)
512+
func (p PlanetScaleEdgeDatabase) printQueryResult(qr *sqltypes.Result, tableNamespace, tableName string, ps *PlanetScaleSource) {
513+
data := QueryResultToRecords(qr, ps)
514514

515515
for _, record := range data {
516516
p.Logger.Record(tableNamespace, tableName, record)

cmd/internal/types.go

Lines changed: 81 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -137,7 +137,7 @@ func TableCursorToSerializedCursor(cursor *psdbconnect.TableCursor) (*Serialized
137137
return sc, nil
138138
}
139139

140-
func QueryResultToRecords(qr *sqltypes.Result) []map[string]interface{} {
140+
func QueryResultToRecords(qr *sqltypes.Result, ps *PlanetScaleSource) []map[string]interface{} {
141141
data := make([]map[string]interface{}, 0, len(qr.Rows))
142142
columns := make([]string, 0, len(qr.Fields))
143143
for _, field := range qr.Fields {
@@ -148,7 +148,14 @@ func QueryResultToRecords(qr *sqltypes.Result) []map[string]interface{} {
148148
record := make(map[string]interface{})
149149
for idx, val := range row {
150150
if idx < len(columns) {
151-
record[columns[idx]] = parseValue(val, qr.Fields[idx].GetColumnType(), qr.Fields[idx].GetType())
151+
parsedValue := parseValue(val, qr.Fields[idx].GetColumnType(), qr.Fields[idx].GetType(), ps)
152+
if parsedValue.isBool {
153+
record[columns[idx]] = parsedValue.boolValue
154+
} else if parsedValue.isNull {
155+
record[columns[idx]] = nil
156+
} else {
157+
record[columns[idx]] = parsedValue.sqlValue
158+
}
152159
}
153160
}
154161
data = append(data, record)
@@ -157,21 +164,60 @@ func QueryResultToRecords(qr *sqltypes.Result) []map[string]interface{} {
157164
return data
158165
}
159166

167+
type Value struct {
168+
sqlValue sqltypes.Value
169+
boolValue bool
170+
isBool bool
171+
isNull bool
172+
}
173+
160174
// After the initial COPY phase, enum and set values may appear as an index instead of a value.
161175
// For example, a value might look like a "1" instead of "apple" in an enum('apple','banana','orange') column)
162-
func parseValue(val sqltypes.Value, columnType string, queryColumnType query.Type) sqltypes.Value {
176+
func parseValue(val sqltypes.Value, columnType string, queryColumnType query.Type, ps *PlanetScaleSource) Value {
177+
if val.IsNull() {
178+
return Value{
179+
isNull: true,
180+
}
181+
}
182+
163183
switch queryColumnType {
164184
case query.Type_DATETIME, query.Type_DATE, query.Type_TIME:
165185
return formatISO8601(queryColumnType, val)
166186
case query.Type_ENUM:
167187
values := parseEnumOrSetValues(columnType)
168-
return mapEnumValue(val, values)
188+
return Value{
189+
sqlValue: mapEnumValue(val, values),
190+
}
169191
case query.Type_SET:
170192
values := parseEnumOrSetValues(columnType)
171-
return mapSetValue(val, values)
193+
return Value{
194+
sqlValue: mapSetValue(val, values),
195+
}
196+
}
197+
198+
if strings.ToLower(columnType) == "tinyint(1)" && !ps.Options.DoNotTreatTinyIntAsBoolean {
199+
return mapTinyIntToBool(val)
172200
}
173201

174-
return val
202+
return Value{
203+
sqlValue: val,
204+
}
205+
}
206+
207+
func mapTinyIntToBool(val sqltypes.Value) Value {
208+
sqlVal, err := val.ToBool()
209+
210+
// Fallback to the original value if we can't convert to bool
211+
if err != nil {
212+
return Value{
213+
sqlValue: val,
214+
}
215+
}
216+
217+
return Value{
218+
boolValue: sqlVal,
219+
isBool: true,
220+
}
175221
}
176222

177223
// Takes enum or set column type like ENUM('a','b','c') or SET('a','b','c')
@@ -190,9 +236,7 @@ func parseEnumOrSetValues(columnType string) []string {
190236
return values
191237
}
192238

193-
func formatISO8601(mysqlType query.Type, value sqltypes.Value) sqltypes.Value {
194-
parsedDatetime := value.ToString()
195-
239+
func formatISO8601(mysqlType query.Type, value sqltypes.Value) Value {
196240
var formatString string
197241
var layout string
198242
if mysqlType == query.Type_DATE {
@@ -202,14 +246,36 @@ func formatISO8601(mysqlType query.Type, value sqltypes.Value) sqltypes.Value {
202246
formatString = "2006-01-02 15:04:05"
203247
layout = time.RFC3339
204248
}
205-
mysqlTime, err := time.Parse(formatString, parsedDatetime)
206-
if err != nil {
207-
// fallback to default value if datetime is not parseable
208-
return value
249+
250+
var (
251+
mysqlTime time.Time
252+
err error
253+
)
254+
255+
if !value.IsNull() {
256+
parsedDatetime := value.ToString()
257+
// Check for zero date
258+
if parsedDatetime == "0000-00-00 00:00:00" || parsedDatetime == "0000-00-00" {
259+
// Use zero epoch time to represent non-null zero date
260+
mysqlTime = time.Unix(0, 0).UTC()
261+
} else {
262+
mysqlTime, err = time.Parse(formatString, parsedDatetime)
263+
if err != nil {
264+
// fallback to default value if datetime is not parseable
265+
return Value{
266+
sqlValue: value,
267+
}
268+
}
269+
}
270+
209271
}
272+
210273
iso8601Datetime := mysqlTime.Format(layout)
211-
formattedValue, _ := sqltypes.NewValue(value.Type(), []byte(iso8601Datetime))
212-
return formattedValue
274+
formattedValue, _ := sqltypes.NewValue(mysqlType, []byte(iso8601Datetime))
275+
276+
return Value{
277+
sqlValue: formattedValue,
278+
}
213279
}
214280

215281
func mapSetValue(value sqltypes.Value, values []string) sqltypes.Value {

cmd/internal/types_test.go

Lines changed: 65 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -99,7 +99,7 @@ func TestCanMapEnumAndSetValues(t *testing.T) {
9999
},
100100
}
101101

102-
output := QueryResultToRecords(&input)
102+
output := QueryResultToRecords(&input, &PlanetScaleSource{})
103103
assert.Equal(t, 2, len(output))
104104
firstRow := output[0]
105105
assert.Equal(t, "active", firstRow["status"].(sqltypes.Value).ToString())
@@ -109,13 +109,65 @@ func TestCanMapEnumAndSetValues(t *testing.T) {
109109
assert.Equal(t, "San Francisco,Oakland", secondRow["locations"].(sqltypes.Value).ToString())
110110
}
111111

112+
func TestCanMapTinyIntValues(t *testing.T) {
113+
input := sqltypes.Result{
114+
Fields: []*query.Field{
115+
{Name: "verified", Type: query.Type_INT8, ColumnType: "tinyint(1)"},
116+
},
117+
Rows: [][]sqltypes.Value{
118+
{sqltypes.NewInt8(1)},
119+
{sqltypes.NewInt8(0)},
120+
},
121+
}
122+
123+
output := QueryResultToRecords(&input, &PlanetScaleSource{
124+
Options: CustomSourceOptions{
125+
DoNotTreatTinyIntAsBoolean: false,
126+
},
127+
})
128+
129+
assert.Equal(t, 2, len(output))
130+
firstRow := output[0]
131+
assert.Equal(t, true, firstRow["verified"].(bool))
132+
secondRow := output[1]
133+
assert.Equal(t, false, secondRow["verified"].(bool))
134+
135+
input = sqltypes.Result{
136+
Fields: []*query.Field{
137+
{Name: "verified", Type: query.Type_INT8, ColumnType: "tinyint(1)"},
138+
},
139+
Rows: [][]sqltypes.Value{
140+
{sqltypes.NewInt8(1)},
141+
{sqltypes.NewInt8(0)},
142+
},
143+
}
144+
145+
output = QueryResultToRecords(&input, &PlanetScaleSource{
146+
Options: CustomSourceOptions{
147+
DoNotTreatTinyIntAsBoolean: true,
148+
},
149+
})
150+
151+
assert.Equal(t, 2, len(output))
152+
firstRow = output[0]
153+
assert.Equal(t, sqltypes.NewInt8(1), firstRow["verified"])
154+
secondRow = output[1]
155+
assert.Equal(t, sqltypes.NewInt8(0), secondRow["verified"])
156+
}
157+
112158
func TestCanFormatISO8601Values(t *testing.T) {
113159
datetimeValue, err := sqltypes.NewValue(query.Type_DATETIME, []byte("2025-02-14 08:08:08"))
114160
assert.NoError(t, err)
115161
dateValue, err := sqltypes.NewValue(query.Type_DATE, []byte("2025-02-14"))
116162
assert.NoError(t, err)
117163
timestampValue, err := sqltypes.NewValue(query.Type_TIMESTAMP, []byte("2025-02-14 08:08:08"))
118164
assert.NoError(t, err)
165+
zeroDatetimeValue, err := sqltypes.NewValue(query.Type_DATETIME, []byte("0000-00-00 00:00:00"))
166+
assert.NoError(t, err)
167+
zeroDateValue, err := sqltypes.NewValue(query.Type_DATE, []byte("0000-00-00"))
168+
assert.NoError(t, err)
169+
zeroTimestampValue, err := sqltypes.NewValue(query.Type_TIMESTAMP, []byte("0000-00-00 00:00:00"))
170+
assert.NoError(t, err)
119171
input := sqltypes.Result{
120172
Fields: []*query.Field{
121173
{Name: "datetime_created_at", Type: sqltypes.Datetime, ColumnType: "datetime"},
@@ -124,13 +176,23 @@ func TestCanFormatISO8601Values(t *testing.T) {
124176
},
125177
Rows: [][]sqltypes.Value{
126178
{datetimeValue, dateValue, timestampValue},
179+
{sqltypes.NULL, sqltypes.NULL, sqltypes.NULL},
180+
{zeroDatetimeValue, zeroDateValue, zeroTimestampValue},
127181
},
128182
}
129183

130-
output := QueryResultToRecords(&input)
131-
assert.Equal(t, 1, len(output))
184+
output := QueryResultToRecords(&input, &PlanetScaleSource{})
185+
assert.Equal(t, 3, len(output))
132186
row := output[0]
133187
assert.Equal(t, "2025-02-14T08:08:08Z", row["datetime_created_at"].(sqltypes.Value).ToString())
134188
assert.Equal(t, "2025-02-14", row["date_created_at"].(sqltypes.Value).ToString())
135189
assert.Equal(t, "2025-02-14T08:08:08Z", row["timestamp_created_at"].(sqltypes.Value).ToString())
190+
nullRow := output[1]
191+
assert.Equal(t, nil, nullRow["datetime_created_at"])
192+
assert.Equal(t, nil, nullRow["date_created_at"])
193+
assert.Equal(t, nil, nullRow["timestamp_created_at"])
194+
zeroRow := output[2]
195+
assert.Equal(t, "1970-01-01T00:00:00Z", zeroRow["datetime_created_at"].(sqltypes.Value).ToString())
196+
assert.Equal(t, "1970-01-01", zeroRow["date_created_at"].(sqltypes.Value).ToString())
197+
assert.Equal(t, "1970-01-01T00:00:00Z", zeroRow["timestamp_created_at"].(sqltypes.Value).ToString())
136198
}

0 commit comments

Comments
 (0)