Skip to content

Commit 524eaca

Browse files
committed
add type check
1 parent 3870816 commit 524eaca

8 files changed

Lines changed: 404 additions & 19 deletions

‎src/client/README.mbt.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -91,7 +91,7 @@ async fn _transaction_example(client : @client.Client) -> Unit {
9191

9292
## Errors
9393

94-
Driver operations raise `ClientError`.
94+
Driver operations raise `ClientError`. Type mismatches detected before encoding or decoding raise `ClientError::WrongType`.
9595

9696
```mbt check
9797
///|

‎src/client/client.mbt‎

Lines changed: 40 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1452,6 +1452,39 @@ fn execute_portal_bytes(portal_name : BytesView, max_rows : Int) -> Bytes raise
14521452
buf.to_bytes()
14531453
}
14541454

1455+
///|
1456+
priv struct EncodedParam {
1457+
format : Int
1458+
value : Bytes?
1459+
}
1460+
1461+
///|
1462+
fn encode_param(param : &ToSql, type_ : Type) -> EncodedParam raise {
1463+
guard param.accepts(type_) else {
1464+
raise wrong_type_error(param.moonbit_type_name(), type_)
1465+
}
1466+
let payload = @buffer.new()
1467+
let value = match param.to_sql(type_, payload) {
1468+
@proto.IsNull::Yes => None
1469+
@proto.IsNull::No => Some(payload.to_bytes())
1470+
}
1471+
{ format: param.format(type_), value }
1472+
}
1473+
1474+
///|
1475+
fn write_encoded_param(
1476+
param : EncodedParam,
1477+
payload : @buffer.Buffer,
1478+
) -> @proto.IsNull {
1479+
match param.value {
1480+
None => @proto.IsNull::Yes
1481+
Some(bytes) => {
1482+
payload.write_bytes(bytes)
1483+
@proto.IsNull::No
1484+
}
1485+
}
1486+
}
1487+
14551488
///|
14561489
fn write_bind(
14571490
buf : @buffer.Buffer,
@@ -1461,14 +1494,17 @@ fn write_bind(
14611494
params : Array[&ToSql],
14621495
describe_portal : Bool,
14631496
) -> Unit raise {
1464-
let formats = params.mapi((index, param) => param.format(param_types[index]))
1465-
let values = params.mapi((index, param) => (param, param_types[index]))
1497+
let encoded_params : Array[EncodedParam] = []
1498+
for index, param in params {
1499+
encoded_params.push(encode_param(param, param_types[index]))
1500+
}
1501+
let formats = encoded_params.map(param => param.format)
14661502
@frontend.bind(
14671503
portal,
14681504
statement,
14691505
formats.iter(),
1470-
values.iter(),
1471-
(pair, payload) => serialize_param(pair.0, pair.1, payload),
1506+
encoded_params.iter(),
1507+
(param, payload) => write_encoded_param(param, payload),
14721508
[1].iter(),
14731509
buf,
14741510
)
@@ -1691,17 +1727,6 @@ fn join_clauses(clauses : Array[String]) -> String {
16911727
out
16921728
}
16931729

1694-
///|
1695-
fn serialize_param(
1696-
param : &ToSql,
1697-
type_ : Type,
1698-
payload : @buffer.Buffer,
1699-
) -> @proto.IsNull raise @proto.ProtocolError {
1700-
param.to_sql(type_, payload) catch {
1701-
err => raise @proto.ProtocolError::InvalidInput(err.to_string())
1702-
}
1703-
}
1704-
17051730
///|
17061731
async fn connect_stream(config : Config) -> Stream {
17071732
let conn = @socket.Tcp::connect_to_host(config.host, port=config.port)

‎src/client/pkg.generated.mbti‎

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@ pub suberror ClientError {
1818
Ssl(String)
1919
Encode(String)
2020
Decode(String)
21-
WrongType(String)
21+
WrongType(WrongTypeError)
2222
ColumnNotFound(String)
2323
RowCount(String)
2424
UnexpectedMessage(String)
@@ -209,7 +209,7 @@ pub async fn Transaction::commit(Self) -> Unit
209209
pub async fn Transaction::execute(Self, String, params? : Array[&ToSql]) -> Int
210210
pub async fn Transaction::prepare(Self, String) -> Statement
211211
pub async fn Transaction::query(Self, String, params? : Array[&ToSql]) -> RowStream
212-
pub async fn Transaction::query_portal(Self, Portal, Int) -> RowStream
212+
pub fn Transaction::query_portal(Self, Portal, Int) -> RowStream raise
213213
pub async fn Transaction::rollback(Self) -> Unit
214214
pub async fn Transaction::transaction(Self) -> Self
215215

@@ -257,11 +257,18 @@ pub fn Type::uuid_array() -> Self
257257
pub fn Type::varchar() -> Self
258258
pub fn Type::varchar_array() -> Self
259259

260+
pub struct WrongTypeError {
261+
moonbit_type : String
262+
postgres_type : Type
263+
} derive(Eq, Show)
264+
260265
// Type aliases
261266

262267
// Traits
263268
pub(open) trait FromSql {
264269
from_sql(Type, Int, BytesView) -> Self raise
270+
accepts(Type) -> Bool
271+
moonbit_type_name() -> String
265272
from_sql_null(Type, Int) -> Self raise = _
266273
}
267274
pub impl FromSql for Bool
@@ -276,6 +283,8 @@ pub impl FromSql for Bytes
276283

277284
pub(open) trait ToSql {
278285
format(Self, Type) -> Int = _
286+
accepts(Self, Type) -> Bool
287+
moonbit_type_name(Self) -> String
279288
to_sql(Self, Type, @buffer.Buffer) -> @protocol.IsNull raise
280289
}
281290
pub impl ToSql for Bool

0 commit comments

Comments
 (0)