// 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/column/columns.go 的
// `Column.FieldData()` 一族与 client/milvusclient/write_options.go 的
// `processInsertColumns`(Apache-2.0)。
//
// 回读侧 `@column.from_field_data` 走的是 `FieldData` → `Column`;
// 这里走反方向 `Column` → `FieldData`。两边共用 `@column.ColumnValue`,
// 所以列类型只在 `column/` 定义一次 —— 写侧多出来的只有「行数对齐」
// 与「向量序列化」。
///|
/// 写入用的一列。名字加数据,等价于上游 `column.Column` 里写入侧用到的那部分。
///
/// 复用 `@column.ColumnValue` 而不是另造一套枚举,是为了让回读的
/// `Column` 能被原样改一改名字写回去(改列名重插入这种用法)。
pub(all) struct WriteColumn {
name : String
value : @column.ColumnValue
}
///|
pub fn WriteColumn::new(
name : String,
value : @column.ColumnValue,
) -> WriteColumn {
{ name, value, }
}
///|
/// 这一列的行数。
pub fn WriteColumn::len(self : WriteColumn) -> Int {
match self.value {
Bool(v) => v.length()
Int8(v) | Int16(v) | Int32(v) => v.length()
Int64(v) => v.length()
Float(v) => v.length()
Double(v) => v.length()
String(v) | VarChar(v) | Text(v) => v.length()
Timestamptz(v) => v.length()
Json(v) => v.length()
Geometry(v) => v.length()
Array(v) => v.length()
FloatVector(_, rows) | Float16Vector(_, rows) | BFloat16Vector(_, rows) =>
rows.length()
BinaryVector(_, rows) | Int8Vector(_, rows) => rows.length()
SparseFloatVector(rows) => rows.length()
}
}
///|
/// 这一列的 `@schema.DataType`。
pub fn WriteColumn::data_type(self : WriteColumn) -> @entity.DataType {
match self.value {
Bool(_) => @entity.DataType::Bool
Int8(_) => @entity.DataType::Int8
Int16(_) => @entity.DataType::Int16
Int32(_) => @entity.DataType::Int32
Int64(_) => @entity.DataType::Int64
Float(_) => @entity.DataType::Float
Double(_) => @entity.DataType::Double
String(_) => @entity.DataType::String
VarChar(_) => @entity.DataType::VarChar
Text(_) => @entity.DataType::Text
Timestamptz(_) => @entity.DataType::Timestamptz
Json(_) => @entity.DataType::Json
Geometry(_) => @entity.DataType::Geometry
Array(_) => @entity.DataType::Array
FloatVector(_, _) => @entity.DataType::FloatVector
Float16Vector(_, _) => @entity.DataType::Float16Vector
BFloat16Vector(_, _) => @entity.DataType::BFloat16Vector
BinaryVector(_, _) => @entity.DataType::BinaryVector
Int8Vector(_, _) => @entity.DataType::Int8Vector
SparseFloatVector(_) => @entity.DataType::SparseFloatVector
}
}
///|
/// 一列 → 一条 `FieldData`。
///
/// 动态字段(JSON)要带上 `is_dynamic`,否则服务端会把它当成 schema 里
/// 不存在的列拒绝。这里靠名字判断:`@column` 侧回读动态字段时也是用
/// `FieldData.is_dynamic` 标记的,写侧没有别的信号可用。
///
/// `nullable` 为真时写一条全 1 的 `valid_data`:调用方给的是实体数组,
/// 目前没有「某行是 null」的表示,写成全 1 与「全部有值」等价,
/// 但显式声明可空位图后服务端不会把空列当坏数据。
pub fn WriteColumn::to_field_data(
self : WriteColumn,
is_dynamic? : Bool = false,
) -> @schema.FieldData raise ClientError {
let rows = self.len()
let field = match self.value {
Bool(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::BoolData(@schema.BoolArray::BoolArray(v)),
),
)
Int8(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::IntData(
@schema.IntArray::IntArray(v.map(narrow_int)),
),
),
)
Int16(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::IntData(
@schema.IntArray::IntArray(v.map(narrow_int)),
),
),
)
Int32(v) =>
@schema.FieldData_Field::Scalars(
scalar(@schema.ScalarField_Data::IntData(@schema.IntArray::IntArray(v))),
)
Int64(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::LongData(@schema.LongArray::LongArray(v)),
),
)
Float(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::FloatData(@schema.FloatArray::FloatArray(v)),
),
)
Double(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::DoubleData(
@schema.DoubleArray::DoubleArray(v),
),
),
)
String(v) | VarChar(v) | Text(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::StringData(
@schema.StringArray::StringArray(v),
),
),
)
// JSON 在 wire 上是 `bytes`,每行一段 UTF-8 文本。
Json(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::JsonData(@schema.JSONArray::JSONArray(v)),
),
)
Timestamptz(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::TimestamptzData(
@schema.TimestamptzArray::TimestamptzArray(v),
),
),
)
Geometry(v) =>
@schema.FieldData_Field::Scalars(
scalar(
@schema.ScalarField_Data::GeometryWktData(
@schema.GeometryWktArray::GeometryWktArray(v),
),
),
)
Array(_) =>
raise ClientError::Encode(
"column \{self.name}: Array fields need an element type, which WriteColumn cannot infer",
)
FloatVector(dim, rows) =>
vector_row_payload(self.name, dim, rows.length(), () => {
@schema.VectorField_Data::FloatVector(
@schema.FloatArray::FloatArray(flatten(rows)),
)
})
Float16Vector(dim, rows) =>
vector_row_payload(self.name, dim, rows.length(), () => {
@schema.VectorField_Data::Float16Vector(
bytes_of_rows(rows, @entity.float16_vector_bytes),
)
})
BFloat16Vector(dim, rows) =>
vector_row_payload(self.name, dim, rows.length(), () => {
@schema.VectorField_Data::Bfloat16Vector(
bytes_of_rows(rows, @entity.bfloat16_vector_bytes),
)
})
BinaryVector(dim, rows) =>
vector_row_payload(self.name, dim, rows.length(), () => {
@schema.VectorField_Data::BinaryVector(join_bytes(rows))
})
Int8Vector(dim, rows) =>
vector_row_payload(self.name, dim, rows.length(), () => {
@schema.VectorField_Data::Int8Vector(join_bytes(rows))
})
SparseFloatVector(rows) =>
@schema.FieldData_Field::Vectors(@schema.VectorField::{
dim: 0L,
valid_data: [],
data: @schema.VectorField_Data::SparseFloatVector(@schema.SparseFloatArray::{
contents: rows.map(row => {
row.to_bytes() catch {
err => raise ClientError::Schema(entity_schema_message(err))
}
}),
dim: max_sparse_dim(rows).to_int64(),
}),
})
}
@schema.FieldData::{
type_: data_type_to_proto(self.data_type()),
field_name: self.name,
field_id: 0L,
is_dynamic,
valid_data: [],
field,
}
|> ignore_row_count(rows)
}
///|
/// 行数为 0 的列没有可写的内容,但也不该静默通过:
/// 服务端拿到空 `FieldData` 会把整批算成 0 行,调用方多半是搞错了。
fn ignore_row_count(
field : @schema.FieldData,
rows : Int,
) -> @schema.FieldData raise ClientError {
if rows == 0 {
raise ClientError::Encode(
"column \{field.field_name} has no rows; every column in a write must carry at least one",
)
}
field
}
///|
/// 标量 payload:只填 data,`valid_data` 留空表示「全部有值」。
fn scalar(data : @schema.ScalarField_Data) -> @schema.ScalarField {
@schema.ScalarField::{ valid_data: [], data, }
}
///|
/// 向量 payload 的公共外壳,顺带校一下 dim 与行数。
fn vector_row_payload(
name : String,
dim : Int,
rows : Int,
build : () -> @schema.VectorField_Data raise ClientError,
) -> @schema.FieldData_Field raise ClientError {
if dim <= 0 {
raise ClientError::Encode(
"column \{name}: vector dim must be positive, got \{dim}",
)
}
if rows == 0 {
raise ClientError::Encode("column \{name}: vector column has no rows")
}
@schema.FieldData_Field::Vectors(@schema.VectorField::{
dim: dim.to_int64(),
valid_data: [],
data: build(),
})
}
///|
/// 逐行编码后拼成一整块字节。fp16 / bf16 在 wire 上是「所有行首尾相接」,
/// 不是每行一个 bytes,所以这里必须拼平。
fn bytes_of_rows(
rows : Array[Array[Float]],
encode : (Array[Float]) -> Bytes,
) -> Bytes {
join_bytes(rows.map(encode))
}
///|
fn join_bytes(parts : Array[Bytes]) -> Bytes {
let buf = Buffer()
for part in parts {
buf.write_bytes(part)
}
buf.to_bytes()
}
///|
fn flatten(rows : Array[Array[Float]]) -> Array[Float] {
let values : Array[Float] = []
for row in rows {
values.append(row)
}
values
}
///|
/// 稀疏向量那一列的 `dim` 是「本批次里最大的下标 + 1」,
/// 与上游 `SparseFloatArray.dim` 的语义一致(注释写明是最大维度)。
fn max_sparse_dim(rows : Array[@entity.SparseFloatVector]) -> UInt {
let mut max : UInt = 0
for row in rows {
for index in row.indices {
if index >= max {
max = index + 1U
}
}
}
max
}
///|
/// `Int` → `Int32` 的窄化。越界的值直接报错,
/// 不静默按低位截断 —— 写进去一个错的值比写失败更难查。
fn narrow_int(value : Int) -> Int raise ClientError {
if value < -2147483648 || value > 2147483647 {
raise ClientError::Encode(
"value \{value} does not fit in a 32-bit integer field",
)
}
value
}