// Licensed to the LF AI & Data foundation under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//
// 移植自 milvus-io/milvus client/milvusclient/{client.go,collection.go,
// write.go,read.go}(Apache-2.0)。
//
// 上游的 `Client` 是一层薄客户端:Option → protobuf 请求 → 发 RPC →
// 响应反序列化,本地不做任何向量计算。这里保持同样的形状。
///|
/// 一次 unary 调用的形状:给路径和已编码的 protobuf body,拿回已解码前的
/// 响应字节。
///
/// 把它抽成函数值而不是直接依赖 `@transport/native`,是因为模块
/// `preferred_target = "wasm"`:`client` 这层必须在 wasm / js 下编得过,
/// 而真连 socket 的实现是 native 专属的(见 `client/native`)。
pub type Unary = async (String, Bytes) -> Bytes raise @transport.RpcError
///|
/// 连到 Milvus 的客户端。
///
/// `config` 只用来取 dbName、超时、metadata 之类的默认值;真正把字
/// 节发出去的是 `call`。
pub struct Client {
config : @transport.Config
call : Unary
}
///|
/// 一次调用失败。四种来源分开,是因为上层的处置不一样:
/// 传输问题可以重试,服务端拒绝要看错误码,schema 不合说明本地模型过时,
/// 编解码失败说明客户端有 bug。
pub suberror ClientError {
Transport(@transport.RpcError)
Server(@errors.MerError)
Schema(String)
Encode(String)
Decode(String)
}
///|
/// 传输层是否还能用。`Client` 的方法里,只有这一类才值得原样重试。
pub fn ClientError::is_transport(self : ClientError) -> Bool {
match self {
Transport(_) => true
_ => false
}
}
///|
/// 服务端是否把这次失败标成可重试。
/// 与 `@errors.is_retryable_err` 一致:只看服务端下发的 `retriable`,
/// 不做本地猜测。
pub fn ClientError::is_retryable(self : ClientError) -> Bool {
match self {
Transport(err) => err.is_retryable()
Server(err) => @errors.is_retryable_err(Some(err))
_ => false
}
}
///|
/// 这次失败是不是超时。
pub fn ClientError::is_deadline_exceeded(self : ClientError) -> Bool {
match self {
Transport(err) => err.is_deadline_exceeded()
_ => false
}
}
///|
/// 拿服务端错误码;不是服务端拒绝就返回 `@errors.unexpected_code`。
pub fn ClientError::code(self : ClientError) -> Int {
match self {
Server(err) => @errors.error_code(Some(err))
_ => @errors.unexpected_code
}
}
///|
/// 底层错误对象,便于上层用 `@errors.same_code` 之类的判等。
pub fn ClientError::server_error(self : ClientError) -> @errors.MerError? {
match self {
Server(err) => Some(err)
_ => None
}
}
///|
pub impl Show for ClientError with fn to_string(self) {
match self {
Transport(err) => "transport: " + err.to_string()
Server(err) => "server: " + err.to_string()
Schema(msg) => "schema: " + msg
Encode(msg) => "encode: " + msg
Decode(msg) => "decode: " + msg
}
}
///|
/// 挂上一条已有的连接。`call` 通常是 `@native.Client::unary` 的部分应用。
pub fn Client::new(config : @transport.Config, call : Unary) -> Client {
{ config, call, }
}
///|
/// 客户端持有的配置。
pub fn Client::config(self : Client) -> @transport.Config {
self.config
}
///|
/// 把服务端 `common.Status` 翻成 `ClientError`。
///
/// Milvus 的老式 RPC 把错误放在响应体的 `status` 里而不是 gRPC 状态里,
/// 所以每个响应都要过一道这个检查。OK 返回 `None`。
fn check_status(status : @common.Status?) -> ClientError? {
if @errors.is_ok(status) {
return None
}
match @errors.from_status(status.unwrap()) {
Some(err) => Some(Server(err))
None => Some(Server(@errors.err_unexpected()))
}
}
///|
/// 发一次调用并解码响应。
///
/// 这里的泛型约束是生成物给的:请求要能 `Write` + `Sized`(编码),
/// 响应要能 `Read` + `Default`(解码)。
async fn[Req : @lib.Write, Resp : @lib.Read + Default] Client::call_service(
self : Client,
path : String,
request : Req,
) -> Resp raise ClientError {
let body = encode_request(request)
let response = (self.call)(path, body) catch {
err => raise ClientError::Transport(err)
}
decode_response(response)
}
///|
fn[Req : @lib.Write] encode_request(request : Req) -> Bytes raise ClientError {
let buf = @buffer.Buffer::Buffer()
@lib.Write::write(request, buf) catch {
err => raise ClientError::Encode(err.to_string())
}
buf.to_bytes()
}
///|
fn[M : @lib.Read + Default] decode_response(
bytes : Bytes,
) -> M raise ClientError {
if bytes.length() == 0 {
// 空 body 也是合法的「默认消息」,gRPC 允许;生成物的 Default 正好接住。
return Default::default()
}
let reader = @lib.BytesReader::from_bytes(bytes)
@lib.Read::read_with_limit(@lib.LimitedReader::LimitedReader(reader)) catch {
err => raise ClientError::Decode(err.to_string())
}
}
///|
/// `@column.ColumnError` 的文本。它在 `@column` 里只 derive 了 `Debug`,
/// 包外拿不到内部字符串,所以这里按构造器解出来。
pub fn column_error_message(err : @column.ColumnError) -> String {
match err {
DataTypeNotMatch(msg) => msg
UnsupportedType(msg) => msg
IndexOutOfRange(msg) => msg
MalformedPayload(msg) => msg
NullValue(msg) => msg
}
}
///|
/// `@column.ColumnValue` 的类型名。`@column` 没有把它暴露出来,
/// 而包外的报错文本需要一个能定位的说法。
pub fn column_value_type_name(value : @column.ColumnValue) -> String {
match value {
Bool(_) => "Bool"
Int8(_) => "Int8"
Int16(_) => "Int16"
Int32(_) => "Int32"
Int64(_) => "Int64"
Float(_) => "Float"
Double(_) => "Double"
String(_) => "String"
VarChar(_) => "VarChar"
Text(_) => "Text"
Timestamptz(_) => "Timestamptz"
Json(_) => "Json"
Geometry(_) => "Geometry"
Array(_) => "Array"
FloatVector(_, _) => "FloatVector"
Float16Vector(_, _) => "Float16Vector"
BFloat16Vector(_, _) => "BFloat16Vector"
BinaryVector(_, _) => "BinaryVector"
Int8Vector(_, _) => "Int8Vector"
SparseFloatVector(_) => "SparseFloatVector"
}
}