diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..48dbffd --- /dev/null +++ b/.dockerignore @@ -0,0 +1,39 @@ +# Rust build artifacts +**/target/ +**/Cargo.lock + +# Python +**/__pycache__/ +**/*.pyc +**/*.pyo +**/*.pyd +.Python +**/.venv/ +**/venv/ +**/.pytest_cache/ +**/.coverage +**/*.egg-info/ + +# Node +**/node_modules/ +**/npm-debug.log +**/yarn-error.log + +# IDEs +.idea/ +.vscode/ +*.swp +*.swo +*~ + +# OS +.DS_Store +Thumbs.db + +# Git +.git/ +.gitignore + +# Misc +*.log +.env diff --git a/README.md b/README.md index b9707ef..b2cdc51 100644 --- a/README.md +++ b/README.md @@ -97,6 +97,16 @@ _ROOST (Robust Open Online Safety Tools), a non-profit organization that brings docker compose up -d ``` + or using the wrapper script + + ```bash + ./start.sh + ``` + + this starts the osprey-worker on its own along with all its required dependencies. + + alternatively, you can start Osprey with `osprey-coordinator`, refer to the [Coordinator README](./example_docker_compose/run_osprey_with_coordinator/README.md) for more information + 6. (Optional) **Port Forward the UI/UI API:** If you are running the docker compose on a headless machine, you will need to port forward the UI and UI API. diff --git a/docker-compose.yaml b/docker-compose.yaml index 4f165c8..cbb553d 100644 --- a/docker-compose.yaml +++ b/docker-compose.yaml @@ -258,7 +258,7 @@ services: ports: - "5432:5432" volumes: - - metadata_data:/var/lib/postgresql + - metadata_data:/var/lib/postgresql/data environment: - POSTGRES_PASSWORD=FoolishPassword - POSTGRES_USER=osprey @@ -377,44 +377,3 @@ services: - ./druid/specs:/specs command: ["/bin/sh", "/specs/submit-specs.sh"] restart: "no" - - osprey_coordinator: - container_name: osprey_coordinator - build: - context: . - dockerfile: osprey_coordinator/Dockerfile - ports: - - "19950:19950" - - "19951:19951" - environment: - - RUST_LOG=info - - ETCD_PEERS=http://etcd:2379 - - SNOWFLAKE_API_ENDPOINT=http://snowflake-id-worker:8088 - - POD_IP=127.0.0.1 - - PUBSUB_EMULATOR_HOST=127.0.0.1:8085 - depends_on: - etcd: - condition: service_healthy - snowflake-id-worker: - condition: service_started - - etcd: - image: quay.io/coreos/etcd:v3.5.15 - ports: - - "2379:2379" - environment: - - ETCD_ENABLE_V2=true - - ETCD_LISTEN_CLIENT_URLS=http://0.0.0.0:2379 - - ETCD_ADVERTISE_CLIENT_URLS=http://0.0.0.0:2379 - healthcheck: - test: - [ - "CMD", - "etcdctl", - "--endpoints=http://localhost:2379", - "endpoint", - "health", - ] - interval: 5s - timeout: 3s - retries: 3 diff --git a/example_data/generate_coordinator_test_data.sh b/example_data/generate_coordinator_test_data.sh new file mode 100755 index 0000000..6e34c0c --- /dev/null +++ b/example_data/generate_coordinator_test_data.sh @@ -0,0 +1,119 @@ +#!/bin/bash + +# continuously generate and send test actions to the osprey coordinator via gRPC +# this mimics the Kafka test data generator but sends directly to the coordinator + +set -e + +COORDINATOR_HOST="${COORDINATOR_HOST:-localhost:19951}" + +# Check if grpcurl is installed +if ! command -v grpcurl &> /dev/null; then + echo "Error: grpcurl is not installed." + echo "Install it with: brew install grpcurl (macOS) or go install github.com/fullstorydev/grpcurl/cmd/grpcurl@latest" + exit 1 +fi + +# Check if jq is installed +if ! command -v jq &> /dev/null; then + echo "Error: jq is not installed." + echo "Install it with: brew install jq (macOS)" + exit 1 +fi + +# Initialize action_id counter +action_id=1 + +# Words to randomly generate post content +words=(hello the quick brown fox jumps over lazy dog and cat runs fast) + +# Get script directory +SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd ) + +# Function to generate random user ID +generate_random_user_id() { + echo "user_$(shuf -i 100-9999 -n 1)" +} + +# Function to generate current timestamp in RFC3339 format +generate_timestamp() { + date -u +"%Y-%m-%dT%H:%M:%S.000000000Z" +} + +# Function to generate random post text +generate_random_text() { + echo "${words[RANDOM % ${#words[@]}]} ${words[RANDOM % ${#words[@]}]} ${words[RANDOM % ${#words[@]}]} ${words[RANDOM % ${#words[@]}]} ${words[RANDOM % ${#words[@]}]}." +} + +# Function to generate action data from template +generate_action() { + local text=$(generate_random_text) + local timestamp=$(generate_timestamp) + local user_id=$(generate_random_user_id) + local ip_address="192.168.1.$(shuf -i 1-254 -n 1)" + + local sed_commands=() + sed_commands+=("s/\$text/$text/g") + sed_commands+=("s/\$timestamp/$timestamp/g") + sed_commands+=("s/\$user_id/$user_id/g") + sed_commands+=("s/\$ip_address/$ip_address/g") + sed_commands+=("s/\$action_id/$action_id/g") + + # Apply all sed commands to template.json + local cmd="sed" + for sed_cmd in "${sed_commands[@]}"; do + cmd="$cmd -e '$sed_cmd'" + done + eval "$cmd" "$SCRIPT_DIR/template.json" +} + +# Function to send a single action +send_action() { + local kafka_format_json=$(generate_action) + + # Extract the data object from the Kafka format and convert to coordinator format + local action_data=$(echo "$kafka_format_json" | jq -c '.data') + local timestamp=$(echo "$kafka_format_json" | jq -r '.send_time') + local action_name=$(echo "$action_data" | jq -r '.action_name') + local data_payload=$(echo "$action_data" | jq -c '.data') + + echo "[$action_id] Sending action - Name: $action_name, Timestamp: $timestamp" + + # Build gRPC request format + jq -n \ + --arg action_id "$action_id" \ + --arg action_name "$action_name" \ + --argjson data_payload "$data_payload" \ + --arg timestamp "$timestamp" \ + '{ + action_id: ($action_id | tonumber), + action_name: $action_name, + action_data_json: ($data_payload | tostring), + timestamp: $timestamp + }' | grpcurl -plaintext -d @ "$COORDINATOR_HOST" \ + osprey.rpc.osprey_coordinator.sync_action.v1.OspreyCoordinatorSyncActionService/ProcessAction + + # Increment action_id + ((action_id++)) +} + +# Function to handle cleanup on script termination +cleanup() { + echo + echo "Stopping data generation..." + exit 0 +} + +# Set up signal handlers for graceful shutdown +trap cleanup SIGINT SIGTERM + +# Main execution +echo "Generating actions every second to Osprey Coordinator at $COORDINATOR_HOST" +echo "Press Ctrl+C to stop..." +echo + +# Infinite loop to generate and send actions +while true; do + send_action + sleep 1 +done diff --git a/example_docker_compose/run_osprey_with_coordinator/README.md b/example_docker_compose/run_osprey_with_coordinator/README.md new file mode 100644 index 0000000..75d9360 --- /dev/null +++ b/example_docker_compose/run_osprey_with_coordinator/README.md @@ -0,0 +1,195 @@ +While Osprey worker can stand on its own by directly ingesting data from Kafka, Osprey Coordinator provides an alternative that provides additional features such as load balancing and synchronous actions. + +## Quick Start + +The easiest way to run Osprey with the Coordinator is using the helper script from the repository root: + +```bash +# Start with coordinator +./start.sh --with-coordinator + +# Start in detached mode +./start.sh --with-coordinator up -d + +# Start with test data producer +./start.sh --with-coordinator --profile coordinator_test_data up +``` + +Or manually using docker compose override files: + +```bash +# From the repository root +docker compose -f docker-compose.yaml -f example_docker_compose/run_osprey_with_coordinator/docker-compose.coordinator.yaml up +``` + +## Overview + +The **Osprey Coordinator** is a Rust-based service that acts as a central hub for distributing actions to Osprey Workers for rule evaluation. It provides two primary modes for receiving actions: + +1. **Bidirectional gRPC Streaming** - Workers connect to the coordinator via persistent bidirectional streams +2. **Synchronous gRPC API** - External services send actions directly for immediate processing + +The coordinator can consume actions from Kafka, Pubsub and/or receives them via gRPC, manages action distribution across connected workers, handles acknowledgments, and ensures reliable action processing. + +## Architecture + +### Components + +``` +┌───────────────────────────────┐ +│ Kafka Topics and/or Pubsub │ +│ (actions_input) │ +└──────────┬────────────────────┘ + │ + ▼ +┌─────────────────────────────────────────────┐ +│ Osprey Coordinator (Rust) │ +│ ┌─────────────────────────────────────┐ │ +│ │ Priority Queue │ │ +│ │ - Sync Actions (high priority) │ │ +│ │ - Async Actions (lower priority) │ │ +│ └─────────────────────────────────────┘ │ +│ │ +│ gRPC Services: │ +│ - Bidirectional Stream (port 19950) │ +│ - Sync Action API (port 19951) │ +└──────────────┬──────────────────────────────┘ + │ + ▼ + ┌──────────────────────┐ + │ Osprey Workers │ + │ (Python) │ + │ - Process rules │ + │ - Send verdicts │ + └──────────────────────┘ +``` + +## Configuration + +The coordinator is configured via the `docker-compose.coordinator.yaml` override file in this directory (`example_docker_compose/run_osprey_with_coordinator/`). This file adds the coordinator service and modifies the worker configuration to connect to it. + +### Environment Variables + +Configure the coordinator via environment variables in `docker-compose.yaml`: + +| Variable | Default | Description | +|----------|---------|-------------| +| `OSPREY_COORDINATOR_BIDI_STREAM_PORT` | `19950` | Port for bidirectional streaming | +| `OSPREY_COORDINATOR_SYNC_ACTION_PORT` | `19951` | Port for synchronous action API | +| `SNOWFLAKE_API_ENDPOINT` | `http://snowflake-id-worker:8088` | Snowflake ID service endpoint | +| `ETCD_PEERS` | `http://etcd:2379` | etcd connection string | +| `OSPREY_KAFKA_BOOTSTRAP_SERVERS` | `kafka:29092` | Kafka broker addresses | +| `OSPREY_KAFKA_INPUT_STREAM_TOPIC` | `osprey.actions_input` | Kafka topic to consume | +| `OSPREY_KAFKA_GROUP_ID` | `osprey_coordinator_group` | Kafka consumer group ID | +| `OSPREY_COORDINATOR_CONSUMER_TYPE` | `kafka` | Consumer type: `kafka` or `pubsub` | +| `MAX_TIME_TO_SEND_TO_ASYNC_QUEUE_MS` | `500` | Max time to wait before queuing async actions | +| `MAX_ACKING_RECEIVER_WAIT_TIME_MS` | `60000` | Max time to wait for worker ack/nack | + +### Example Configuration + +To customize the coordinator, edit `example_docker_compose/run_osprey_with_coordinator/docker-compose.coordinator.yaml`: + +```yaml +services: + osprey-coordinator: + environment: + - RUST_LOG=info + - ETCD_PEERS=http://etcd:2379 + - SNOWFLAKE_API_ENDPOINT=http://snowflake-id-worker:8088 + - OSPREY_COORDINATOR_CONSUMER_TYPE=kafka # or 'pubsub' + - OSPREY_COORDINATOR_BIDI_STREAM_PORT=19950 + - OSPREY_COORDINATOR_SYNC_ACTION_PORT=19951 +``` + +## Using the Coordinator + +**Worker Configuration:** + +### Worker Connection + +When using `docker-compose.coordinator.yaml`, workers are automatically configured to connect to the coordinator. The override file sets: + +```yaml +osprey-worker: + environment: + - OSPREY_INPUT_STREAM_SOURCE=osprey_coordinator + - OSPREY_COORDINATOR_SERVICE_NAME=osprey_coordinator +``` + +**How It Works:** + +1. Worker connects to coordinator on port 19950 +2. Worker sends initial connection request with client ID +3. Coordinator sends actions to worker via the bidirectional stream +4. Worker processes actions through rules +5. Worker sends ack/nack with optional verdicts back to coordinator +6. Connection automatically reconnects every 60-120 seconds for load balancing + + +### Direct Action Submission (Sync API) + +External services can submit actions directly to the coordinator for synchronous processing. + +**Using grpcurl:** + +```bash +# Send a single action for immediate processing +grpcurl -plaintext \ + -d '{ + "action_id": 12345, + "action_name": "user_login", + "action_data_json": "{\"user_id\":\"user_123\",\"ip_address\":\"192.168.1.1\"}", + "timestamp": "2024-11-25T10:30:00.000000000Z" + }' \ + localhost:19951 \ + osprey.rpc.osprey_coordinator.sync_action.v1.OspreyCoordinatorSyncActionService/ProcessAction +``` + +or Use the test data producer: + +```bash +./start.sh --with-coordinator --profile coordinator_test_data up +``` + +### Kafka/Pubsub Integration +The coordinator can automatically consume from either Kafka or PubSub (but not both simultaneously). Set `OSPREY_COORDINATOR_CONSUMER_TYPE` to choose: + +- `kafka` (default) - Consume from Kafka +- `pubsub` - Consume from Google Cloud PubSub + +Configure the appropriate environment variables for your chosen consumer in `example_docker_compose/run_osprey_with_coordinator/docker-compose.coordinator.yaml`: + +```yaml + osprey-coordinator: + environment: + # Consumer selection (kafka or pubsub) + - OSPREY_COORDINATOR_CONSUMER_TYPE=kafka + # Kafka configuration (when using kafka) + - OSPREY_KAFKA_BOOTSTRAP_SERVERS=kafka:29092 + - OSPREY_KAFKA_INPUT_STREAM_TOPIC=osprey.actions_input + - OSPREY_KAFKA_GROUP_ID=osprey_coordinator_group + + # Pubsub + - OSPREY_COORDINATOR_SERVICE_ACCOUNT + - PUBSUB_SUBSCRIPTION_PROJECT_ID + - PUBSUB_SUBSCRIPTION_ID + - PUBSUB_ENCRYPTION_KEY_URI + # Optionally + - PUBSUB_MAX_MESSAGES = 5000 # default + - PUBSUB_MAX_PROCESSING_MESSAGES = 5000 # default + + # shared by both Kafka and Pubsub, optional + - MAX_TIME_TO_SEND_TO_ASYNC_QUEUE_MS = 500 # default + - MAX_ACKING_RECEIVER_WAIT_TIME_MS = 6000 # default +``` + + +**Sending Actions via Kafka:** + +Use the test data producer: + +```bash +# Start the Kafka test data producer +./start.sh --with-coordinator --profile test_data up kafka-test-data-producer -d +``` + diff --git a/example_docker_compose/run_osprey_with_coordinator/docker-compose.coordinator.yaml b/example_docker_compose/run_osprey_with_coordinator/docker-compose.coordinator.yaml new file mode 100644 index 0000000..6a898dc --- /dev/null +++ b/example_docker_compose/run_osprey_with_coordinator/docker-compose.coordinator.yaml @@ -0,0 +1,109 @@ +# Docker Compose override for running Osprey with Coordinator +# Use with: docker compose -f docker-compose.yaml -f docker-compose.coordinator.yaml up +# Or use the helper script: ./start.sh --with-coordinator + +services: + # Override worker to connect to coordinator instead of Kafka directly + osprey-worker: + depends_on: + kafka: + condition: service_healthy + kafka-topic-creator: + condition: service_completed_successfully + bigtable: + condition: service_healthy + bigtable-initializer: + condition: service_completed_successfully + minio: + condition: service_healthy + minio-bucket-init: + condition: service_completed_successfully + etcd: + condition: service_healthy + postgres: + condition: service_healthy + osprey-coordinator: + condition: service_started + environment: + - ETCD_PEERS=http://etcd:2379 + - OSPREY_INPUT_STREAM_SOURCE=osprey_coordinator + - OSPREY_COORDINATOR_SERVICE_NAME=osprey_coordinator + + # Add Osprey Coordinator service + osprey-coordinator: + container_name: osprey-coordinator + hostname: osprey-coordinator + build: + context: . + dockerfile: osprey_coordinator/Dockerfile + ports: + - "19950:19950" + - "19951:19951" + environment: + - RUST_LOG=info + - ETCD_PEERS=http://etcd:2379 + - SNOWFLAKE_API_ENDPOINT=http://snowflake-id-worker:8088 + - OSPREY_COORDINATOR_CONSUMER_TYPE=kafka + - OSPREY_COORDINATOR_BIDI_STREAM_PORT=19950 + - OSPREY_COORDINATOR_SYNC_ACTION_PORT=19951 + - POD_IP=osprey-coordinator + - OSPREY_KAFKA_BOOTSTRAP_SERVERS=kafka:29092 + - OSPREY_KAFKA_INPUT_STREAM_TOPIC=osprey.actions_input + - OSPREY_KAFKA_GROUP_ID=osprey_coordinator_group + - MAX_TIME_TO_SEND_TO_ASYNC_QUEUE_MS=500 + - MAX_ACKING_RECEIVER_WAIT_TIME_MS=60000 + depends_on: + etcd: + condition: service_healthy + snowflake-id-worker: + condition: service_started + kafka: + condition: service_healthy + kafka-topic-creator: + condition: service_completed_successfully + + # Add etcd service (required by coordinator) + etcd: + image: quay.io/coreos/etcd:v3.5.15 + ports: + - "2379:2379" + environment: + - ETCD_ENABLE_V2=true + - ETCD_LISTEN_CLIENT_URLS=http://0.0.0.0:2379 + - ETCD_ADVERTISE_CLIENT_URLS=http://0.0.0.0:2379 + healthcheck: + test: + [ + "CMD", + "etcdctl", + "--endpoints=http://localhost:2379", + "endpoint", + "health", + ] + interval: 5s + timeout: 3s + retries: 3 + + # Optional coordinator test data generator + # Run with: docker compose -f docker-compose.yaml -f docker-compose.coordinator.yaml --profile coordinator_test_data up coordinator-test-data-producer + coordinator-test-data-producer: + image: alpine:latest + hostname: coordinator-test-data-producer + container_name: coordinator-test-data-producer + depends_on: + osprey-coordinator: + condition: service_started + profiles: + - coordinator_test_data + - coordinator-test-data + environment: + COORDINATOR_HOST: "osprey-coordinator:19951" + volumes: + - ./example_data:/osprey/example_data + command: + - /bin/sh + - -c + - | + apk add --no-cache bash jq curl && + wget -qO- https://github.com/fullstorydev/grpcurl/releases/download/v1.8.7/grpcurl_1.8.7_linux_x86_64.tar.gz | tar xz -C /usr/local/bin && + /osprey/example_data/generate_coordinator_test_data.sh diff --git a/osprey_coordinator/Cargo.toml b/osprey_coordinator/Cargo.toml index 742cbd3..510678e 100644 --- a/osprey_coordinator/Cargo.toml +++ b/osprey_coordinator/Cargo.toml @@ -77,6 +77,7 @@ tracing-subscriber = "0.2" trust-dns-resolver = "0.21" uuid = { version = "1.0", features = ["v4"] } which = "4.4.0" +rdkafka = "0.38.0" [build-dependencies] glob = "0.3" diff --git a/osprey_coordinator/src/consumer/kafka.rs b/osprey_coordinator/src/consumer/kafka.rs new file mode 100644 index 0000000..3314347 --- /dev/null +++ b/osprey_coordinator/src/consumer/kafka.rs @@ -0,0 +1,350 @@ +use crate::consumer::message_consumer::{ConsumerConfig, ConsumerMessage, MessageConsumer}; +use crate::consumer::message_decoder; +use crate::coordinator_metrics::OspreyCoordinatorMetrics; +use crate::metrics::counters::StaticCounter; +use crate::metrics::histograms::StaticHistogram; +use crate::priority_queue::{AckOrNack, AckableAction, PriorityQueueSender}; +use crate::signals::exit_signal; +use crate::snowflake_client::SnowflakeClient; +use anyhow::Result; +use async_trait::async_trait; +use prost_types::Timestamp; +use rdkafka::config::ClientConfig; +use rdkafka::consumer::{Consumer, StreamConsumer}; +use rdkafka::error::KafkaError; +use rdkafka::message::{Headers, Message as KafkaRawMessage}; +use std::collections::HashMap; +use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; +use tokio::time::{timeout, Instant}; + +pub struct KafkaConsumer { + consumer: StreamConsumer, +} + +pub struct KafkaMessage { + data: Vec, + attributes: HashMap, + timestamp: Timestamp, + id: String, +} + +impl KafkaMessage { + pub fn new( + data: Vec, + attributes: HashMap, + timestamp: Timestamp, + id: String, + ) -> Self { + Self { + data, + attributes, + timestamp, + id, + } + } +} + +impl ConsumerMessage for KafkaMessage { + fn data(&self) -> &[u8] { + &self.data + } + + fn attributes(&self) -> &HashMap { + &self.attributes + } + + fn timestamp(&self) -> Timestamp { + self.timestamp.clone() + } + + fn id(&self) -> String { + self.id.clone() + } +} + +#[async_trait] +impl MessageConsumer for KafkaConsumer { + type Message = KafkaMessage; + type Error = KafkaError; + + async fn receive(&mut self) -> Result { + let msg = self.consumer.recv().await?; + + let data = msg.payload().unwrap_or(&[]).to_vec(); + + let attributes: HashMap = msg + .headers() + .map(|headers| { + headers + .iter() + .filter_map(|header| { + let key = header.key.to_string(); + let value = header + .value + .and_then(|v| String::from_utf8(v.to_vec()).ok()) + .unwrap_or_default(); + Some((key, value)) + }) + .collect() + }) + .unwrap_or_default(); + + let timestamp_millis = msg.timestamp().to_millis().unwrap_or_else(|| { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_millis() as i64 + }); + + let timestamp = Timestamp { + seconds: timestamp_millis / 1000, + nanos: ((timestamp_millis % 1000) * 1_000_000) as i32, + }; + + let partition = msg.partition(); + let offset = msg.offset(); + let id = format!("kafka-{}-{}", partition, offset); + + Ok(KafkaMessage::new(data, attributes, timestamp, id)) + } + + async fn ack(&self, _message: &Self::Message) -> Result<(), Self::Error> { + self.consumer + .commit_consumer_state(rdkafka::consumer::CommitMode::Async)?; + Ok(()) + } + + async fn nack(&self, _message: &Self::Message) -> Result<(), Self::Error> { + Ok(()) + } +} + +impl KafkaConsumer { + pub async fn new() -> Result { + let input_topic = std::env::var("OSPREY_KAFKA_INPUT_STREAM_TOPIC") + .unwrap_or("osprey.actions_input".to_string()); + let input_bootstrap_servers = + std::env::var("OSPREY_KAFKA_BOOTSTRAP_SERVERS").unwrap_or("localhost:9092".to_string()); + let group_id = std::env::var("OSPREY_KAFKA_GROUP_ID") + .unwrap_or("osprey_coordinator_group".to_string()); + + tracing::info!( + "Creating Kafka consumer for topic: {} with bootstrap servers: {}", + input_topic, + input_bootstrap_servers + ); + + let consumer: StreamConsumer = ClientConfig::new() + .set("group.id", &group_id) + .set("bootstrap.servers", &input_bootstrap_servers) + .set("enable.auto.commit", "false") + .set("auto.offset.reset", "earliest") + .create::()?; + + consumer.subscribe(&[&input_topic])?; + + Ok(Self { consumer }) + } +} + +pub async fn start_kafka_consumer( + snowflake_client: Arc, + priority_queue_sender: PriorityQueueSender, + metrics: Arc, +) -> Result<()> { + tracing::info!("Kafka consumer starting..."); + + let mut consumer = KafkaConsumer::new().await?; + let config = ConsumerConfig::default(); + + loop { + tokio::select! { + _ = exit_signal() => { + tracing::info!("Received exit signal, shutting down Kafka consumer"); + return Ok(()); + } + message_result = consumer.receive() => { + let message = match message_result { + Ok(msg) => msg, + Err(e) => { + tracing::error!({error = %e}, "[kafka] error receiving message"); + continue; + } + }; + + let ack_id: u64 = rand::Rng::gen(&mut rand::thread_rng()); + let message_id = message.id(); + + let action = match message.attributes().get("encoding").map(|s| s.as_str()) { + Some("proto") => { + message_decoder::decode_proto_message( + message.data(), + ack_id, + message.timestamp(), + &snowflake_client, + &metrics, + ) + .await + } + _ => { + message_decoder::decode_msgpack_json_message( + message.data(), + ack_id, + message.timestamp(), + &snowflake_client, + &metrics, + ) + .await + } + }; + + let action = match action { + Ok(action) => action, + Err(e) => { + tracing::error!( + {error = %e, ack_id = %ack_id, message_id = %message_id}, + "[kafka] failed to decode message" + ); + if let Err(nack_err) = consumer.nack(&message).await { + tracing::error!( + {error = %nack_err, message_id = %message_id}, + "[kafka] failed to nack message" + ); + } + continue; + } + }; + + let (ackable_action, acking_receiver) = AckableAction::new(action); + + tracing::debug!( + {ack_id = %ack_id, message_id = %message_id}, + "[kafka] received message" + ); + + let send_start_time = Instant::now(); + match timeout( + config.max_time_to_send_to_async_queue, + priority_queue_sender.send_async(ackable_action), + ) + .await + { + Ok(Ok(())) => { + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[kafka] sent message to priority queue" + ); + metrics.async_classification_added_to_queue.incr(); + } + Ok(Err(e)) => { + tracing::error!( + {error = %e, message_id = %message_id}, + "[kafka] priority queue send error" + ); + if let Err(nack_err) = consumer.nack(&message).await { + tracing::error!( + {error = %nack_err, message_id = %message_id}, + "[kafka] failed to nack message" + ); + } + continue; + } + Err(_) => { + tracing::error!( + {message_id = %message_id}, + "[kafka] sending to priority queue timed out" + ); + if let Err(nack_err) = consumer.nack(&message).await { + tracing::error!( + {error = %nack_err, message_id = %message_id}, + "[kafka] failed to nack message" + ); + } + continue; + } + } + metrics + .priority_queue_send_time_async + .record(send_start_time.elapsed()); + + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[kafka] waiting on ack or nack" + ); + + let receive_start_time = Instant::now(); + match timeout(config.max_acking_receiver_wait_time, acking_receiver).await { + Ok(Ok(ack_or_nack)) => match ack_or_nack { + AckOrNack::Ack(_) => { + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[kafka] acking message" + ); + metrics.async_classification_result_ack.incr(); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + + if let Err(e) = consumer.ack(&message).await { + tracing::error!( + {error = %e, message_id = %message_id}, + "[kafka] failed to ack message" + ); + } + } + AckOrNack::Nack => { + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[kafka] nacking message" + ); + metrics.async_classification_result_nack.incr(); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + + if let Err(e) = consumer.nack(&message).await { + tracing::error!( + {error = %e, message_id = %message_id}, + "[kafka] failed to nack message" + ); + } + } + }, + Ok(Err(recv_error)) => { + tracing::error!( + {message_id = %message_id, recv_error = %recv_error, ack_id = %ack_id}, + "[kafka] acking sender dropped" + ); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + + if let Err(e) = consumer.nack(&message).await { + tracing::error!( + {error = %e, message_id = %message_id}, + "[kafka] failed to nack message" + ); + } + } + Err(_) => { + tracing::error!( + {message_id = %message_id, ack_id = %ack_id}, + "[kafka] waiting for ack/nack timed out" + ); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + + if let Err(e) = consumer.nack(&message).await { + tracing::error!( + {error = %e, message_id = %message_id}, + "[kafka] failed to nack message" + ); + } + } + } + } + } + } +} diff --git a/osprey_coordinator/src/consumer/message_consumer.rs b/osprey_coordinator/src/consumer/message_consumer.rs new file mode 100644 index 0000000..08ff132 --- /dev/null +++ b/osprey_coordinator/src/consumer/message_consumer.rs @@ -0,0 +1,49 @@ +use anyhow::Result; +use async_trait::async_trait; +use prost_types::Timestamp; +use std::collections::HashMap; +use tokio::time::Duration as TokioDuration; + +#[derive(Clone)] +pub struct ConsumerConfig { + pub max_time_to_send_to_async_queue: TokioDuration, + pub max_acking_receiver_wait_time: TokioDuration, +} + +impl Default for ConsumerConfig { + fn default() -> Self { + Self { + max_time_to_send_to_async_queue: TokioDuration::from_millis( + std::env::var("MAX_TIME_TO_SEND_TO_ASYNC_QUEUE_MS") + .unwrap_or("500".to_string()) + .parse::() + .unwrap(), + ), + max_acking_receiver_wait_time: TokioDuration::from_millis( + std::env::var("MAX_ACKING_RECEIVER_WAIT_TIME_MS") + .unwrap_or("60000".to_string()) + .parse::() + .unwrap(), + ), + } + } +} + +pub trait ConsumerMessage { + fn data(&self) -> &[u8]; + fn attributes(&self) -> &HashMap; + fn timestamp(&self) -> Timestamp; + fn id(&self) -> String; +} + +#[async_trait] +pub trait MessageConsumer: Send { + type Message: ConsumerMessage; + type Error: std::error::Error + Send + Sync + 'static; + + async fn receive(&mut self) -> Result; + + async fn ack(&self, message: &Self::Message) -> Result<(), Self::Error>; + + async fn nack(&self, message: &Self::Message) -> Result<(), Self::Error>; +} diff --git a/osprey_coordinator/src/consumer/message_decoder.rs b/osprey_coordinator/src/consumer/message_decoder.rs new file mode 100644 index 0000000..cf52a37 --- /dev/null +++ b/osprey_coordinator/src/consumer/message_decoder.rs @@ -0,0 +1,90 @@ +use anyhow::{anyhow, Result}; +use convert_case::{Case, Casing}; +use prost::Message as ProstMessage; +use prost_types::Timestamp; +use serde::Deserialize; +use serde_json::Value; + +use crate::{ + coordinator_metrics::OspreyCoordinatorMetrics, + metrics::counters::StaticCounter, + proto::{ + self, osprey_coordinator_action::ActionData, osprey_coordinator_action::SecretData, + Action as OspreyProtoAction, + }, + snowflake_client::SnowflakeClient, +}; + +pub async fn decode_proto_message( + message_data: &[u8], + ack_id: u64, + message_timestamp: Timestamp, + snowflake_client: &SnowflakeClient, + metrics: &OspreyCoordinatorMetrics, +) -> Result { + let osprey_proto_action = OspreyProtoAction::decode(message_data)?; + let action_id = if osprey_proto_action.id == 0 { + metrics.action_id_snowflake_generation_proto.incr(); + snowflake_client.generate_id().await? + } else { + osprey_proto_action.id + }; + let action_name = osprey_proto_action + .data + .ok_or_else(|| anyhow!("missing action data"))? + .to_string() + .to_case(Case::Snake); + Ok(proto::OspreyCoordinatorAction { + ack_id, + action_id, + action_name, + action_data: Some(ActionData::ProtoActionData(message_data.into())), + secret_data: None, + timestamp: Some(message_timestamp), + }) +} + +pub async fn decode_msgpack_json_message( + message_data: &[u8], + ack_id: u64, + message_timestamp: Timestamp, + snowflake_client: &SnowflakeClient, + metrics: &OspreyCoordinatorMetrics, +) -> Result { + use msgpack_simple::MsgPack; + + #[derive(Deserialize, Debug)] + struct MsgpackAction { + id: Option, + name: String, + data: Value, + secret_data: Option, + } + + let decoded = MsgPack::parse(message_data)?; + let decoded = decoded.as_string()?; + let action: MsgpackAction = serde_json::from_str(decoded.as_str())?; + + let serde_json_vec = serde_json::to_vec(&action.data)?; + let optional_secret_data = match &action.secret_data { + Some(secret_data) => Some(SecretData::JsonSecretData(serde_json::to_vec(secret_data)?)), + _ => None, + }; + + let action_id = match action.id { + Some(id) => id.parse::()?, + None => { + metrics.action_id_snowflake_generation_json.incr(); + snowflake_client.generate_id().await? + } + }; + + Ok(proto::OspreyCoordinatorAction { + ack_id, + action_id, + action_name: action.name, + action_data: Some(ActionData::JsonActionData(serde_json_vec)), + secret_data: optional_secret_data, + timestamp: Some(message_timestamp), + }) +} diff --git a/osprey_coordinator/src/consumer/mod.rs b/osprey_coordinator/src/consumer/mod.rs new file mode 100644 index 0000000..519e8c0 --- /dev/null +++ b/osprey_coordinator/src/consumer/mod.rs @@ -0,0 +1,7 @@ +pub mod kafka; +pub mod message_consumer; +pub mod message_decoder; +pub mod pubsub; + +pub use kafka::start_kafka_consumer; +pub use pubsub::start_pubsub_subscriber; diff --git a/osprey_coordinator/src/consumer/pubsub.rs b/osprey_coordinator/src/consumer/pubsub.rs new file mode 100644 index 0000000..9b5331c --- /dev/null +++ b/osprey_coordinator/src/consumer/pubsub.rs @@ -0,0 +1,520 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use crate::consumer::message_consumer::{ConsumerConfig, ConsumerMessage}; +use crate::gcloud::grpc::connection::Connection; +use crate::gcloud::{ + auth::AuthorizationHeaderInterceptor, + gcp_metadata::GCPMetadataClient, + google::pubsub::v1::subscriber_client::SubscriberClient, + kms::{AesGcmEnvelope, GOOGLE_KMS_DOMAIN}, + pubsub::{PubSubSubscription, GOOGLE_PUBSUB_DOMAIN}, +}; +use crate::metrics::counters::StaticCounter; +use crate::metrics::histograms::StaticHistogram; +use crate::metrics::MetricsClientBuilder; +use crate::{ + consumer::message_decoder, + coordinator_metrics::OspreyCoordinatorMetrics, + priority_queue::{AckOrNack, AckableAction, PriorityQueueSender}, + proto, + pub_sub_streaming_pull::DetachedMessage, + pub_sub_streaming_pull::{FlowControl, SpawnTaskPerMessageHandler, StreamingPullManager}, +}; +use anyhow::{anyhow, Result}; +use prost_types::Timestamp; +use rand::Rng; +use tokio::time::{timeout, Instant}; +use tonic::{codegen::InterceptedService, transport::Channel}; + +use crate::signals::exit_signal; +use crate::snowflake_client::SnowflakeClient; + +pub struct PubSubMessage { + inner: DetachedMessage, +} + +impl ConsumerMessage for PubSubMessage { + fn data(&self) -> &[u8] { + &self.inner.data + } + + fn attributes(&self) -> &HashMap { + self.inner.attributes() + } + + fn timestamp(&self) -> Timestamp { + self.inner.publish_time() + } + + fn id(&self) -> String { + self.inner.message_id.clone() + } +} + +impl From for PubSubMessage { + fn from(msg: DetachedMessage) -> Self { + PubSubMessage { inner: msg } + } +} + +async fn decrypt_pubsub_message( + kms_envelope: Arc, + message_data: &[u8], +) -> Result> { + kms_envelope + .decrypt(message_data) + .await + .map_err(|err| anyhow!("message decryption failed: {}", err.to_string())) +} + +async fn create_action_from_pubsub_message( + kms_envelope: Arc, + message_data: &[u8], + message_attributes: &HashMap, + ack_id: u64, + message_timestamp: Timestamp, + snowflake_client: &SnowflakeClient, + metrics: &OspreyCoordinatorMetrics, +) -> Result { + let decrypted_message_vector = match message_attributes.get("encrypted") { + Some(is_encrypted) if is_encrypted == "true" => { + Some(decrypt_pubsub_message(kms_envelope, message_data).await?) + } + _ => None, + }; + let message_data = match &decrypted_message_vector { + Some(data) => &data[..], + None => message_data, + }; + + match message_attributes.get("encoding") { + Some(encoding) if encoding == "proto" => { + message_decoder::decode_proto_message( + message_data, + ack_id, + message_timestamp, + snowflake_client, + metrics, + ) + .await + } + _ => { + message_decoder::decode_msgpack_json_message( + message_data, + ack_id, + message_timestamp, + snowflake_client, + metrics, + ) + .await + } + } +} + +async fn create_pubsub_subscription_client( +) -> SubscriberClient> { + let emulator_host = std::env::var("PUBSUB_EMULATOR_HOST").ok(); + + let timeout = Duration::from_secs(5); + + if let Some(emulator_host) = emulator_host { + tracing::info!("Creating subscription client to emulator"); + Connection::new_no_auth( + format!("http://{}", emulator_host).try_into().unwrap(), + timeout, + ) + .create_subscriber_client() + } else { + tracing::info!("Creating subscription client to real pubsub"); + let service_account = + std::env::var("OSPREY_COORDINATOR_SERVICE_ACCOUNT").unwrap_or("default".to_string()); + let client = GCPMetadataClient::new(service_account).unwrap(); + Connection::from_metadata_client( + client, + timeout, + Duration::from_secs(24000), + GOOGLE_PUBSUB_DOMAIN, + ) + .await + .unwrap() + .create_subscriber_client() + } +} + +pub async fn start_pubsub_subscriber( + snowflake_client: Arc, + priority_queue_sender: PriorityQueueSender, + metrics: Arc, +) -> Result<()> { + let subscriber_client = create_pubsub_subscription_client().await; + let subscription_name = { + let project_id = + std::env::var("PUBSUB_SUBSCRIPTION_PROJECT_ID").unwrap_or("osprey-dev".to_string()); + + let subscription_id = std::env::var("PUBSUB_SUBSCRIPTION_ID") + .unwrap_or("osprey-coordinator-actions".to_string()); + + PubSubSubscription::new(project_id, subscription_id) + }; + + let kek_uri = std::env::var("PUBSUB_ENCRYPTION_KEY_URI").unwrap_or("".to_string()); + + let kms_envelope = Connection::from_metadata_client( + GCPMetadataClient::new("default".into())?, + Duration::from_secs(5), + Duration::from_secs(24000), + GOOGLE_KMS_DOMAIN, + ) + .await? + .create_kms_aes_gcm_envelope(kek_uri, Vec::new(), true)?; + + let kms_envelope = Arc::new(kms_envelope); + let max_messages = std::env::var("PUBSUB_MAX_MESSAGES") + .unwrap_or("5000".to_string()) + .parse::() + .unwrap(); + let max_processing_messages = std::env::var("PUBSUB_MAX_PROCESSING_MESSAGES") + .unwrap_or("5000".to_string()) + .parse::() + .unwrap(); + + let config = ConsumerConfig::default(); + let max_time_to_send_to_async_queue = config.max_time_to_send_to_async_queue; + let max_acking_receiver_wait_time = config.max_acking_receiver_wait_time; + + tracing::info!( + {subscription_name = %subscription_name}, + "creating streaming pull manager" + ); + let flow_control = FlowControl::default() + .set_max_messages(max_messages) + .set_max_processing_messages(max_processing_messages) + .set_max_bytes(1024 * 1024 * 1024); + StreamingPullManager::new( + subscriber_client, + subscription_name, + flow_control, + SpawnTaskPerMessageHandler::new(move |message: DetachedMessage| { + let metrics = metrics.clone(); + let pubsub_message = PubSubMessage::from(message); + let message_id = pubsub_message.id(); + let priority_queue_sender = priority_queue_sender.clone(); + let snowflake_client = snowflake_client.clone(); + let kms_envelope = kms_envelope.clone(); + + async move { + let ack_id: u64 = rand::thread_rng().gen(); + + let action = create_action_from_pubsub_message( + kms_envelope, + pubsub_message.data(), + pubsub_message.attributes(), + ack_id, + pubsub_message.timestamp(), + snowflake_client.as_ref(), + &metrics, + ) + .await + .map_err(|_| ())?; + + let (ackable_action, acking_receiver) = AckableAction::new(action); + + tracing::debug!( + {ack_id = %ack_id, message_id = %message_id}, + "[pubsub] received message" + ); + + let send_start_time = Instant::now(); + match timeout( + max_time_to_send_to_async_queue, + priority_queue_sender.send_async(ackable_action), + ) + .await + { + Ok(Ok(())) => { + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[pubsub] sent message to priority queue" + ); + metrics.async_classification_added_to_queue.incr(); + } + Ok(Err(e)) => { + tracing::error!( + {error = %e, message_id = %message_id}, + "[pubsub] priority queue send error" + ); + return Err(()); + } + Err(_) => { + tracing::error!( + {message_id = %message_id}, + "[pubsub] sending to priority queue timed out" + ); + return Err(()); + } + } + metrics + .priority_queue_send_time_async + .record(send_start_time.elapsed()); + + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[pubsub] waiting on ack or nack" + ); + + let receive_start_time = Instant::now(); + match timeout(max_acking_receiver_wait_time, acking_receiver).await { + Ok(Ok(ack_or_nack)) => match ack_or_nack { + AckOrNack::Ack(_) => { + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[pubsub] acking message" + ); + metrics.async_classification_result_ack.incr(); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + Ok(()) + } + AckOrNack::Nack => { + tracing::debug!( + {message_id = %message_id, ack_id = %ack_id}, + "[pubsub] nacking message" + ); + metrics.async_classification_result_nack.incr(); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + Err(()) + } + }, + Ok(Err(recv_error)) => { + tracing::error!( + {message_id = %message_id, recv_error = %recv_error, ack_id = %ack_id}, + "[pubsub] acking sender dropped" + ); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + Err(()) + } + Err(_) => { + tracing::error!( + {message_id = %message_id, ack_id = %ack_id}, + "[pubsub] waiting for ack/nack timed out" + ); + metrics + .receiver_ack_time_async + .record(receive_start_time.elapsed()); + Err(()) + } + } + } + }), + MetricsClientBuilder::new("osprey_coordinator.pull"), + ) + .gracefully_stop_on_signal(exit_signal(), Duration::from_secs(30)) + .await; + Result::Ok(()) +} + +#[cfg(test)] +mod tests { + use base64::Engine; + use prost_types::Timestamp; + use serde_json::json; + use std::collections::HashMap; + use std::sync::Arc; + + use super::create_action_from_pubsub_message; + use crate::coordinator_metrics::OspreyCoordinatorMetrics; + use crate::proto; + use crate::snowflake_client::SnowflakeClient; + + #[tokio::test] + async fn test_create_action_from_pubsub_message_1() { + use crate::gcloud::grpc::connection::Connection; + use msgpack_simple::MsgPack; + + let action_json = json!({ + "id": "123456789", + "name": "guild_invite_created", + "data": { + "char": "abc", + "int": 1i64, + "float2": 1.1_f64 + }, + }); + let encoded = MsgPack::String(action_json.to_string()).encode(); + + let snowflake = SnowflakeClient::new("http://localhost:8088".to_string()); + let metrics = OspreyCoordinatorMetrics::new(); + // Create a mock KMS envelope (won't be used since encrypted != true) + let connection = Connection::new_no_auth( + "http://localhost:8080".try_into().unwrap(), + std::time::Duration::from_secs(5), + ); + let kms_envelope = Arc::new( + connection + .create_kms_aes_gcm_envelope("gcp-kms://test".to_string(), Vec::new(), false) + .unwrap(), + ); + + let attributes = HashMap::new(); + + let prost_action = create_action_from_pubsub_message( + kms_envelope, + encoded.as_slice(), + &attributes, + 12344242, + Timestamp::default(), + &snowflake, + &metrics, + ) + .await; + + println!("{:?}", prost_action); + assert!(prost_action.is_ok(), "prost action decoding failed"); + let prost_action = prost_action.unwrap(); + assert_eq!(prost_action.action_name, "guild_invite_created"); + + let action_data = prost_action + .action_data + .expect("action_data should be present"); + let proto::osprey_coordinator_action::ActionData::JsonActionData(json_bytes) = action_data + else { + panic!() + }; + let data: serde_json::Value = serde_json::from_slice(&json_bytes).unwrap(); + assert_eq!(data["char"], "abc"); + assert_eq!(data["int"], 1); + assert_eq!(data["float2"], 1.1); + } + + #[tokio::test] + async fn test_create_action_from_pubsub_message_2() { + use crate::gcloud::grpc::connection::Connection; + use msgpack_simple::MsgPack; + + let action_json = json!({ + "id": "123456789", + "name": "guild_invite_created", + "data": { + "char": "abc", + "int": 1i64, + "float2": 1.1_f64 + }, + }); + + let encoded = MsgPack::String(action_json.to_string()).encode(); + + let snowflake = SnowflakeClient::new("http://localhost:8088".to_string()); + let metrics = OspreyCoordinatorMetrics::new(); + let connection = Connection::new_no_auth( + "http://localhost:8080".try_into().unwrap(), + std::time::Duration::from_secs(5), + ); + let kms_envelope = Arc::new( + connection + .create_kms_aes_gcm_envelope("gcp-kms://test".to_string(), Vec::new(), false) + .unwrap(), + ); + let attributes = HashMap::new(); + + let prost_action = create_action_from_pubsub_message( + kms_envelope, + encoded.as_slice(), + &attributes, + 12344242, + Timestamp::default(), + &snowflake, + &metrics, + ) + .await; + + println!("{:?}", prost_action); + assert!(prost_action.is_ok(), "prost action decoding failed"); + let prost_action = prost_action.unwrap(); + assert_eq!(prost_action.action_name, "guild_invite_created"); + + let proto::osprey_coordinator_action::ActionData::JsonActionData(json_bytes) = prost_action + .action_data + .expect("action_data should be present") + else { + panic!("Expected JsonActionData variant") + }; + + let data: serde_json::Value = serde_json::from_slice(&json_bytes).unwrap(); + assert_eq!(data["char"], "abc"); + assert_eq!(data["int"], 1); + assert_eq!(data["float2"], 1.1); + } + + #[tokio::test] + async fn test_create_action_from_pubsub_proto_action() { + use crate::gcloud::grpc::connection::Connection; + use std::fs::File; + use std::io::Read; + + let mut file = File::open("test_data/pubsub_proto_message.json").unwrap(); + let mut data = String::new(); + file.read_to_string(&mut data).unwrap(); + let json: serde_json::Value = serde_json::from_str(&data).unwrap(); + let action_jsons = json.as_array().expect("was not array"); + let action_json = action_jsons[0].as_object().expect("is not map"); + + let action_data_str = action_json + .get("message") + .unwrap() + .as_object() + .unwrap() + .get("data") + .unwrap() + .as_str() + .unwrap(); + + let action_bytes = base64::engine::general_purpose::STANDARD + .decode(action_data_str) + .unwrap(); + + let snowflake = SnowflakeClient::new("http://localhost:8088".to_string()); + let metrics = OspreyCoordinatorMetrics::new(); + let connection = Connection::new_no_auth( + "http://localhost:8080".try_into().unwrap(), + std::time::Duration::from_secs(5), + ); + let kms_envelope = Arc::new( + connection + .create_kms_aes_gcm_envelope("gcp-kms://test".to_string(), Vec::new(), false) + .unwrap(), + ); + let mut attributes = HashMap::new(); + attributes.insert("encoding".to_string(), "proto".to_string()); + + let prost_action = create_action_from_pubsub_message( + kms_envelope, + action_bytes.as_slice(), + &attributes, + 12344242, + Timestamp::default(), + &snowflake, + &metrics, + ) + .await; + + assert!(prost_action.is_ok(), "proto action decoding failed"); + let prost_action = prost_action.unwrap(); + + // Validate we got ProtoActionData (not JsonActionData) and the bytes match input + let proto::osprey_coordinator_action::ActionData::ProtoActionData(proto_bytes) = + prost_action + .action_data + .expect("action_data should be present") + else { + panic!("Expected ProtoActionData variant") + }; + assert_eq!(proto_bytes, action_bytes); + } +} diff --git a/osprey_coordinator/src/main.rs b/osprey_coordinator/src/main.rs index 1d17246..9beb493 100644 --- a/osprey_coordinator/src/main.rs +++ b/osprey_coordinator/src/main.rs @@ -1,5 +1,6 @@ mod backoff_utils; mod cached_futures; +mod consumer; mod coordinator_metrics; mod discovery; mod etcd; @@ -14,7 +15,6 @@ mod pigeon; mod priority_queue; mod proto; mod pub_sub_streaming_pull; -mod pubsub; mod shutdown_handler; mod signals; mod snowflake_client; @@ -34,8 +34,8 @@ use crate::snowflake_client::SnowflakeClient; use crate::metrics::emit_worker::SpawnEmitWorker; use crate::metrics::new_client; +use consumer::{start_kafka_consumer, start_pubsub_subscriber}; use priority_queue::{create_ackable_action_priority_queue, spawn_priority_queue_metrics_worker}; -use pubsub::start_pubsub_subscriber; use tokio::join; use crate::osprey_bidirectional_stream::OspreyCoordinatorServer; @@ -95,11 +95,46 @@ async fn main() -> Result<()> { metrics.clone(), )); - let pubsub_fut = start_pubsub_subscriber( - snowflake_client, - priority_queue_sender.clone(), - metrics.clone(), - ); + let consumer_type = std::env::var("OSPREY_COORDINATOR_CONSUMER_TYPE").ok(); + + let consumer_fut = match consumer_type.as_deref() { + Some("kafka") => { + tracing::info!("starting Kafka consumer"); + Box::pin(start_kafka_consumer( + snowflake_client.clone(), + priority_queue_sender.clone(), + metrics.clone(), + )) + as std::pin::Pin> + Send>> + } + Some("pubsub") => { + tracing::info!("starting PubSub subscriber"); + Box::pin(start_pubsub_subscriber( + snowflake_client.clone(), + priority_queue_sender.clone(), + metrics.clone(), + )) + as std::pin::Pin> + Send>> + } + Some(invalid) => { + anyhow::bail!( + "invalid OSPREY_COORDINATOR_CONSUMER_TYPE '{}', must be 'kafka' or 'pubsub'", + invalid + ); + } + None => { + tracing::info!( + "OSPREY_COORDINATOR_CONSUMER_TYPE not set, defaulting to Kafka consumer" + ); + Box::pin(start_kafka_consumer( + snowflake_client.clone(), + priority_queue_sender.clone(), + metrics.clone(), + )) + as std::pin::Pin> + Send>> + } + }; + let grpc_bidi_stream_service_fut = pigeon::serve( osprey_coordinator_grpc_bidi_stream_service, "osprey_coordinator", @@ -122,14 +157,14 @@ async fn main() -> Result<()> { priority_queue_receiver.clone(), ); - tracing::info!("starting pubsub listener/bidi stream/sync classification rpc"); - let (pubsub_result, grpc_bidi_stream_service_result, sync_action_service_result) = join!( - pubsub_fut, + tracing::info!("starting consumer/bidi stream/sync classification rpc"); + let (consumer_result, grpc_bidi_stream_service_result, sync_action_service_result) = join!( + consumer_fut, grpc_bidi_stream_service_fut, sync_action_service_fut ); tracing::info!({ - pubsub_result=?pubsub_result, + consumer_result=?consumer_result, bidi_stream_result=?grpc_bidi_stream_service_result, sync_action_result=?sync_action_service_result}, "osprey coordinator terminated"); diff --git a/osprey_coordinator/src/pubsub.rs b/osprey_coordinator/src/pubsub.rs deleted file mode 100644 index 697d392..0000000 --- a/osprey_coordinator/src/pubsub.rs +++ /dev/null @@ -1,494 +0,0 @@ -use std::collections::HashMap; -use std::sync::Arc; -use std::time::Duration; - -use crate::gcloud::grpc::connection::Connection; -use crate::gcloud::{ - auth::AuthorizationHeaderInterceptor, - gcp_metadata::GCPMetadataClient, - google::pubsub::v1::subscriber_client::SubscriberClient, - kms::{AesGcmEnvelope, GOOGLE_KMS_DOMAIN}, - pubsub::{PubSubSubscription, GOOGLE_PUBSUB_DOMAIN}, -}; -use crate::metrics::counters::StaticCounter; -use crate::metrics::histograms::StaticHistogram; -use crate::metrics::MetricsClientBuilder; -use crate::{ - coordinator_metrics::OspreyCoordinatorMetrics, - priority_queue::{AckOrNack, AckableAction, PriorityQueueSender}, - proto::{self, osprey_coordinator_action::SecretData}, - pub_sub_streaming_pull::DetachedMessage, - pub_sub_streaming_pull::{FlowControl, SpawnTaskPerMessageHandler, StreamingPullManager}, -}; -use anyhow::{anyhow, Result}; -use msgpack_simple::MsgPack; -use prost::Message; -use prost_types::Timestamp; -use rand::Rng; -use serde::Deserialize; -use serde_json::Value; -use tokio::time::{timeout, Duration as TokioDuration, Instant}; -use tonic::{codegen::InterceptedService, transport::Channel}; - -use crate::proto::Action as OspreyProtoAction; -use crate::signals::exit_signal; -use crate::snowflake_client::SnowflakeClient; -use convert_case::{Case, Casing}; -use proto::osprey_coordinator_action::ActionData; - -async fn decode_proto_message( - message_data: &[u8], - ack_id: u64, - message_timestamp: Timestamp, - snowflake_client: &SnowflakeClient, - metrics: &OspreyCoordinatorMetrics, -) -> Result { - let osprey_proto_action = OspreyProtoAction::decode(message_data).unwrap(); - let action_id = if osprey_proto_action.id == 0 { - metrics.action_id_snowflake_generation_proto.incr(); - snowflake_client.generate_id().await? - } else { - osprey_proto_action.id - }; - let action_name = osprey_proto_action - .data - .unwrap() - .to_string() - .to_case(Case::Snake); - Ok(proto::OspreyCoordinatorAction { - ack_id, - action_id, - action_name, - action_data: Some(ActionData::ProtoActionData(message_data.into())), - secret_data: None, - timestamp: Some(message_timestamp), - }) -} - -async fn decode_msgpack_json_message( - message_data: &[u8], - ack_id: u64, - message_timestamp: Timestamp, - snowflake_client: &SnowflakeClient, - metrics: &OspreyCoordinatorMetrics, -) -> Result { - // This whole function can probably be optimized way better, but in the interest of time I am leaving - // it in a working state for now. - #[derive(Deserialize, Debug)] - struct PubsubAction { - id: Option, - name: String, - data: Value, - secret_data: Option, - } - - let decoded = MsgPack::parse(message_data)?; - let decoded = decoded.as_string()?; - let pubsub_action: PubsubAction = serde_json::from_str(decoded.as_str())?; - - let serde_json_vec = serde_json::to_vec(&pubsub_action.data)?; - let optional_secret_data = match &pubsub_action.secret_data { - Some(secret_data) => Some(SecretData::JsonSecretData(serde_json::to_vec(secret_data)?)), - _ => None, - }; - - // old msgpack parsing - // let mut out = Vec::with_capacity(1024 * 6); - // let mut de = serde_json::Deserializer::from_slice(serde_json_vec.as_slice()); - // let mut se = rmp_serde::Serializer::new(&mut out); - // serde_transcode::transcode(&mut de, &mut se).unwrap(); - - let action_id = match pubsub_action.id { - Some(id) => id.parse::()?, - None => { - metrics.action_id_snowflake_generation_json.incr(); - snowflake_client.generate_id().await? - } - }; - - Ok(proto::OspreyCoordinatorAction { - ack_id, - action_id, - action_name: pubsub_action.name, - action_data: Some(ActionData::JsonActionData(serde_json_vec)), - secret_data: optional_secret_data, - timestamp: Some(message_timestamp), - }) -} - -async fn decrypt_pubsub_message( - kms_envelope: Arc, - message_data: &[u8], -) -> Result> { - kms_envelope - .decrypt(message_data) - .await - .map_err(|err| anyhow!("message decryption failed: {}", err.to_string())) -} - -async fn create_action_from_pubsub_message( - kms_envelope: Arc, - message_data: &[u8], - message_attributes: &HashMap, - ack_id: u64, - message_timestamp: Timestamp, - snowflake_client: &SnowflakeClient, - metrics: &OspreyCoordinatorMetrics, -) -> Result { - let decrypted_message_vector = match message_attributes.get("encrypted") { - Some(is_encrypted) if is_encrypted == "true" => { - Some(decrypt_pubsub_message(kms_envelope, message_data).await?) - } - _ => None, - }; - let message_data = match &decrypted_message_vector { - Some(data) => &data[..], - None => message_data, - }; - match message_attributes.get("encoding") { - Some(encoding) if encoding == "proto" => { - decode_proto_message( - message_data, - ack_id, - message_timestamp, - snowflake_client, - metrics, - ) - .await - } - _ => { - decode_msgpack_json_message( - message_data, - ack_id, - message_timestamp, - snowflake_client, - metrics, - ) - .await - } - } -} - -async fn create_pubsub_subscription_client( -) -> SubscriberClient> { - let emulator_host = std::env::var("PUBSUB_EMULATOR_HOST").ok(); - - let timeout = Duration::from_secs(5); - - if let Some(emulator_host) = emulator_host { - tracing::info!("Creating subscription client to emulator"); - Connection::new_no_auth( - format!("http://{}", emulator_host).try_into().unwrap(), - timeout, - ) - .create_subscriber_client() - } else { - tracing::info!("Creating subscription client to real pubsub"); - let service_account = - std::env::var("OSPREY_COORDINATOR_SERVICE_ACCOUNT").unwrap_or("default".to_string()); - let client = GCPMetadataClient::new(service_account).unwrap(); - Connection::from_metadata_client( - client, - timeout, - Duration::from_secs(24000), - GOOGLE_PUBSUB_DOMAIN, - ) - .await - .unwrap() - .create_subscriber_client() - } -} - -pub async fn start_pubsub_subscriber( - snowflake_client: Arc, - priority_queue_sender: PriorityQueueSender, - metrics: Arc, -) -> Result<()> { - let subscriber_client = create_pubsub_subscription_client().await; - let subscription_name = { - let project_id = - std::env::var("PUBSUB_SUBSCRIPTION_PROJECT_ID").unwrap_or("osprey-dev".to_string()); - - let subscription_id = std::env::var("PUBSUB_SUBSCRIPTION_ID") - .unwrap_or("osprey-coordinator-actions".to_string()); - - PubSubSubscription::new(project_id, subscription_id) - }; - - let kek_uri = std::env::var("PUBSUB_ENCRYPTION_KEY_URI").unwrap_or("".to_string()); - - let kms_envelope = Connection::from_metadata_client( - GCPMetadataClient::new("default".into())?, - Duration::from_secs(5), - Duration::from_secs(24000), - GOOGLE_KMS_DOMAIN, - ) - .await? - .create_kms_aes_gcm_envelope(kek_uri, Vec::new(), true)?; - - let kms_envelope = Arc::new(kms_envelope); - let max_messages = std::env::var("PUBSUB_MAX_MESSAGES") - .unwrap_or("5000".to_string()) - .parse::() - .unwrap(); - let max_processing_messages = std::env::var("PUBSUB_MAX_PROCESSING_MESSAGES") - .unwrap_or("5000".to_string()) - .parse::() - .unwrap(); - let max_time_to_send_to_async_queue = TokioDuration::from_millis( - std::env::var("MAX_TIME_TO_SEND_TO_ASYNC_QUEUE_MS") - .unwrap_or("500".to_string()) - .parse::() - .unwrap(), - ); - let max_acking_receiver_wait_time = TokioDuration::from_millis( - std::env::var("MAX_ACKING_RECEIVER_WAIT_TIME_MS") - .unwrap_or("60000".to_string()) - .parse::() - .unwrap(), - ); - - tracing::info!( - {subscription_name = %subscription_name}, - "creating streaming pull manager" - ); - let flow_control = FlowControl::default() - .set_max_messages(max_messages) - .set_max_processing_messages(max_processing_messages) - .set_max_bytes(1024 * 1024 * 1024); - StreamingPullManager::new( - subscriber_client, - subscription_name, - flow_control, - SpawnTaskPerMessageHandler::new(move |message: DetachedMessage| { - let metrics = metrics.clone(); - let message_id = message.message_id.clone(); - let priority_queue_sender = priority_queue_sender.clone(); - let snowflake_client = snowflake_client.clone(); - let kms_envelope = kms_envelope.clone(); - async move { - let message_attributes = message.attributes(); - - let ack_id: u64 = { - let mut rng = rand::thread_rng(); - rng.gen() - }; - - let action = create_action_from_pubsub_message( - kms_envelope, - message.data.as_slice(), - message_attributes, - ack_id, - message.publish_time(), - snowflake_client.as_ref(), - &metrics, - ).await - .map_err(|_| ())?; - let (ackable_action, acking_receiver) = AckableAction::new(action); - - tracing::debug!({ack_id = %ack_id, message_id=%message_id}, "[pubsub] received pubsub message"); - let send_start_time = Instant::now(); - match timeout(max_time_to_send_to_async_queue, priority_queue_sender.send_async(ackable_action)).await { - Ok(Ok(())) => { - tracing::debug!({message_id=%message_id, ack_id=ack_id}, "[pubsub] sent pubsub message to priority queue"); - metrics.async_classification_added_to_queue.incr(); - }, - Ok(Err(e)) => { - tracing::error!({error=%e},"[pubsub] priority queue send error"); - }, - Err(_) => { - tracing::error!({message_id=%message_id}, "[pubsub] sending to priority queue timed out"); - } - }; - metrics.priority_queue_send_time_async.record(send_start_time.elapsed()); - tracing::debug!({message_id=%message_id, ack_id=ack_id},"[pubsub] waiting on ack or nack"); - - let receive_start_time = Instant::now(); - match timeout(max_acking_receiver_wait_time, acking_receiver).await { // 5 Minutes to return - Ok(Ok(ack_or_nack)) => match ack_or_nack { - AckOrNack::Ack(_optional_execution_result) => { - tracing::debug!({message_id=%message_id, ack_id=ack_id},"[pubsub] acking message"); - metrics.async_classification_result_ack.incr(); - metrics.receiver_ack_time_async.record(receive_start_time.elapsed()); - Ok(()) - }, - AckOrNack::Nack => { - tracing::debug!({message_id=%message_id, ack_id=ack_id},"[pubsub] nacking message"); - metrics.async_classification_result_nack.incr(); - metrics.receiver_ack_time_async.record(receive_start_time.elapsed()); - Err(()) - }, - }, - Ok(Err(recv_error)) => { - tracing::error!({message_id=%message_id, recv_error=%recv_error, ack_id=ack_id},"[pubsub] acking sender dropped"); - metrics.receiver_ack_time_async.record(receive_start_time.elapsed()); - Err(()) - }, - Err(_) => { - tracing::error!({message_id=%message_id, ack_id=ack_id}, "[pubsub] waiting for ack/nack timed out"); - metrics.receiver_ack_time_async.record(receive_start_time.elapsed()); - Err(()) - }, - } - } - }), - MetricsClientBuilder::new("osprey_coordinator.pull"), - ) - .gracefully_stop_on_signal(exit_signal(), Duration::from_secs(30)) - .await; - Result::Ok(()) -} - -// TODO: Fix these tests - -// #[cfg(test)] -// mod tests { -// use std::{collections::HashMap, fs::File, io::Read}; - -// use msgpack_simple::MsgPack; -// use prost::Message; -// use prost_types::Timestamp; -// // use protobuf_json_mapping; -// use serde_json::json; - -// // use discord_smite_rpc_actions_proto::SmiteCoordinatorAction as SmiteRpcAction; -// use crate::proto; -// use std::io::Cursor; -// use std::str; - -// use super::create_action_from_pubsub_message; - -// // #[test] -// // fn test_create_action_from_pubsub_message_1() { -// // let action_json = json!({ -// // "id": "123456789", -// // "name": "guild_invite_created", -// // "data": { -// // "char": "abc", -// // "int": 1i64, -// // "float2": 1.1_f64 -// // }, - -// // }); -// // let encoded = MsgPack::String(action_json.to_string()).encode(); - -// // let prost_action = -// // create_action_from_pubsub_message(encoded.as_slice(), 12344242, Timestamp::default()); -// // println!("{:?}", prost_action); -// // assert!(prost_action.is_ok(), "prost action decoding failed"); -// // let prost_action = prost_action.unwrap(); -// // let data = MsgPack::parse(&prost_action.action_data).unwrap(); -// // println!("{:?}", data); -// // let mut data: HashMap = data -// // .as_map() -// // .unwrap() -// // .into_iter() -// // .map(|v| (v.key.as_string().unwrap(), v.value)) -// // .collect(); - -// // let x = data.remove("char").unwrap().as_string().unwrap(); - -// // assert_eq!(x, "abc".to_string()); -// // } - -// // #[test] -// // fn test_create_action_from_pubsub_message_2() { -// // let action_json = json!({ -// // "id": "123456789", -// // "name": "guild_invite_created", -// // "data": { -// // "char": "abc", -// // "int": 1i64, -// // "float2": 1.1_f64 -// // }, - -// // }); -// // // let encoded = MsgPack::String(action_json.to_string()).encode(); - -// // let prost_action = create_action_from_pubsub_message( -// // action_json.to_string().bytes().collect(), -// // 12344242, -// // Timestamp::default(), -// // ); -// // println!("{:?}", prost_action); -// // assert!(prost_action.is_ok(), "prost action decoding failed"); -// // let prost_action = prost_action.unwrap(); -// // let data = MsgPack::parse(&prost_action.action_data).unwrap(); -// // println!("{:?}", data); -// // let mut data: HashMap = data -// // .as_map() -// // .unwrap() -// // .into_iter() -// // .map(|v| (v.key.as_string().unwrap(), v.value)) -// // .collect(); - -// // let x = data.remove("char").unwrap().as_string().unwrap(); - -// // assert_eq!(x, "abc".to_string()); -// // } -// #[test] -// fn test_create_action_from_pubsub_proto_action() { -// let mut file = File::open("pubsub_messages.json").unwrap(); -// let mut data = String::new(); -// file.read_to_string(&mut data).unwrap(); -// let json: serde_json::Value = serde_json::from_str(&data).unwrap(); -// let action_jsons = json.as_array().expect("was not array"); -// let action_json = action_jsons[0].as_object().expect("is not map"); -// println!("{:?}", action_json); -// let action_data = action_json -// .get("message") -// .unwrap() -// .as_object() -// .unwrap() -// .get("data") -// .unwrap() -// .as_str() -// .unwrap(); - -// println!("{:?}", action_data); -// let action_data = str::from_utf8(action_data.as_bytes()).unwrap(); -// println!("{:?}", action_data); -// // let action_data = base64::decode(action_data).unwrap(); - -// // let mut action_prost: SmiteRpcAction = -// // SmiteRpcAction::decode(&mut Cursor::new(action_data.to_string())).unwrap(); -// println!("--------"); -// let test_object_data = "CWUQhNrjpmQOGsgBCgsJCwCCONT0dQYQARK4AQogfnvDH9y83BpR5FraKYJSCxTDrQkSb1axRgl9i82pYwgQARoLCI2fiZYGEPeX1CciCwjksNCbBhC8i4FgKgsI5MaGmwYQvIuBYDItCg0KCzk1LjIuMTIuMTY2EgkKB0FuZHJvaWQaEQoPRGlzY29yZCBBbmRyb2lkOi8KDwoNMTc2LjIxOS40Mi4zNBIJCgdBbmRyb2lkGhEKD0Rpc2NvcmQgQW5kcm9pZEILCPqihpsGEOj7p0c="; -// let decoded_base64 = base64::decode(test_object_data).unwrap(); -// let mut test_action = proto::SmiteCoordinatorAction::default(); -// test_action.id = 20; - -// println!("{:?}", test_action); -// let x = test_action.encode_to_vec(); -// println!("{:?}", x); -// println!("--------"); -// println!("{:?}", decoded_base64); -// let mut action_prost: SmiteRpcAction = -// SmiteRpcAction::decode(decoded_base64.as_slice()).unwrap(); -// println!("{:?}", action_prost); - -// let output = protobuf_json_mapping::print_to_string(action_prost); -// println!("{:?}", output); - -// // let prost_action = create_action_from_pubsub_message( -// // action_json.to_string().bytes().collect(), -// // 12344242, -// // Timestamp::default(), -// // ); -// // println!("{:?}", prost_action); -// // assert!(prost_action.is_ok(), "prost action decoding failed"); -// // let prost_action = prost_action.unwrap(); -// // let data = MsgPack::parse(&prost_action.action_data).unwrap(); -// // println!("{:?}", data); -// // let mut data: HashMap = data -// // .as_map() -// // .unwrap() -// // .into_iter() -// // .map(|v| (v.key.as_string().unwrap(), v.value)) -// // .collect(); - -// // let x = data.remove("char").unwrap().as_string().unwrap(); - -// // assert_eq!(x, "abc".to_string()); -// } -// } diff --git a/osprey_coordinator/test_data/pubsub_proto_message.json b/osprey_coordinator/test_data/pubsub_proto_message.json new file mode 100644 index 0000000..71fdc63 --- /dev/null +++ b/osprey_coordinator/test_data/pubsub_proto_message.json @@ -0,0 +1,12 @@ +[ + { + "ackId": "test-ack-id-proto", + "message": { + "attributes": { + "encoding": "proto", + "user_id": "987654321" + }, + "data": "CbFo3joAAAAAMgA=" + } + } +] diff --git a/start.sh b/start.sh new file mode 100755 index 0000000..202b0ae --- /dev/null +++ b/start.sh @@ -0,0 +1,69 @@ +#!/bin/bash +set -e + +# Helper script to start Osprey with different configurations +# Usage: +# ./start.sh # Start with worker directly consuming from Kafka +# ./start.sh --with-coordinator # Start with Osprey Coordinator +# ./start.sh --help # Show this help + +show_help() { + echo "Osprey Startup Helper" + echo "" + echo "Usage: ./start.sh [OPTIONS] [COMPOSE_ARGS...]" + echo "" + echo "Options:" + echo " --with-coordinator Start Osprey with Coordinator (workers connect to coordinator)" + echo " --help, -h Show this help message" + echo "" + echo "Examples:" + echo " ./start.sh # Direct Kafka consumption" + echo " ./start.sh --with-coordinator # With coordinator" + echo " ./start.sh --with-coordinator up -d # With coordinator in detached mode" + echo " ./start.sh --with-coordinator --profile test_data up # With test data producer" + echo "" + echo "When using --with-coordinator, the following services are added:" + echo " - osprey-coordinator: Action distribution and load balancing" + echo " - etcd: Service discovery for coordinator" + echo "" +} + +USE_COORDINATOR=false +COMPOSE_FILES="-f docker-compose.yaml" +COMPOSE_ARGS=() + +# Parse arguments +while [[ $# -gt 0 ]]; do + case $1 in + --with-coordinator) + USE_COORDINATOR=true + shift + ;; + --help|-h) + show_help + exit 0 + ;; + *) + # Pass remaining args to docker compose + COMPOSE_ARGS+=("$1") + shift + ;; + esac +done + +if [ "$USE_COORDINATOR" = true ]; then + echo "Starting Osprey with Coordinator..." + COMPOSE_FILES="$COMPOSE_FILES -f example_docker_compose/run_osprey_with_coordinator/docker-compose.coordinator.yaml" +else + echo "Starting Osprey without Coordiantor (direct Kafka consumption)..." +fi + +# If no compose args provided, default to 'up' +if [ ${#COMPOSE_ARGS[@]} -eq 0 ]; then + COMPOSE_ARGS=("up") +fi + +echo "Running: docker compose $COMPOSE_FILES ${COMPOSE_ARGS[@]}" +echo "" + +exec docker compose $COMPOSE_FILES "${COMPOSE_ARGS[@]}"