use std::sync::Arc;
use async_trait::async_trait;
use dataflow_rs::engine::error::DataflowError;
use dataflow_rs::engine::functions::AsyncFunctionHandler;
use dataflow_rs::engine::task_context::TaskContext;
use dataflow_rs::engine::task_outcome::TaskOutcome;
use futures::TryStreamExt;
use mongodb::bson::{self, Document};
use serde_json::Value;
use super::connector_helpers::{
ConnectorCall, apply_output, is_mongo, require_db_connector, resolve_value, timed_query,
to_connect_error,
};
use super::schema::{FieldKind, FieldSchema};
use crate::connector::ConnectorRegistry;
use crate::connector::mongo_pool::MongoPoolCache;
const NAME: &str = "mongo_read";
pub struct MongoReadHandler {
pub pool_cache: Arc<MongoPoolCache>,
pub registry: Arc<ConnectorRegistry>,
pub max_rows: usize,
}
#[async_trait]
impl AsyncFunctionHandler for MongoReadHandler {
type Input = Value;
async fn execute(
&self,
ctx: &mut TaskContext<'_>,
input: &Value,
) -> dataflow_rs::Result<TaskOutcome> {
let call = ConnectorCall::begin(NAME, input, ctx)?;
let database = call.require_str(input, "database")?;
let collection = call.require_str(input, "collection")?;
let filter_val = input
.get("filter")
.map(|f| resolve_value(f, ctx))
.unwrap_or_else(|| Value::Object(serde_json::Map::new()));
let filter_doc = bson::to_document(&filter_val)
.map_err(|e| DataflowError::Validation(format!("Invalid MongoDB filter: {e}")))?;
call.run(&self.registry, async {
let connector_config = call.resolve(&self.registry, Some("read")).await?;
let db_config = require_db_connector(&connector_config, call.connector)?;
if !is_mongo(&db_config.connection_string) {
return Err(DataflowError::Validation(format!(
"{NAME} requires a MongoDB connector, but '{}' has a non-MongoDB \
connection string (expected a mongodb:// or mongodb+srv:// URL)",
call.connector
)));
}
let client = self
.pool_cache
.get_client(call.connector, db_config)
.await
.map_err(to_connect_error)?;
let coll = client.database(database).collection::<Document>(collection);
let max_rows = self.max_rows;
let docs: Vec<Document> = timed_query(db_config.query_timeout_ms, call.name, async {
let mut cursor = coll.find(filter_doc).await.map_err(|e| e.to_string())?;
let mut docs: Vec<Document> = Vec::new();
while let Some(doc) = cursor.try_next().await.map_err(|e| e.to_string())? {
if docs.len() >= max_rows {
return Err(format!(
"{NAME} result exceeds query.max_limit ({max_rows} \
documents) — add a filter/limit or raise the cap"
));
}
docs.push(doc);
}
Ok(docs)
})
.await?;
let result: Vec<Value> = docs
.iter()
.filter_map(|doc| serde_json::to_value(doc).ok())
.collect();
apply_output(ctx, call.output, Value::Array(result));
Ok(TaskOutcome::Success)
})
.await
}
}
pub(super) const MONGO_READ_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "connector",
description: "Name of the MongoDB connector.",
kind: FieldKind::String,
required: true,
resolvable: false,
alias: None,
},
FieldSchema {
name: "database",
description: "Mongo database name.",
kind: FieldKind::String,
required: true,
resolvable: false,
alias: None,
},
FieldSchema {
name: "collection",
description: "Mongo collection name.",
kind: FieldKind::String,
required: true,
resolvable: false,
alias: None,
},
FieldSchema {
name: "filter",
description: "Mongo find() filter document. Defaults to {}. Accepts {\"var\": \"path\"} to read the value from the message.",
kind: FieldKind::Object,
required: false,
resolvable: true,
alias: None,
},
FieldSchema {
name: "output",
description: "Dotted path where matched documents are written.",
kind: FieldKind::String,
required: false,
resolvable: false,
alias: None,
},
];