from __future__ import annotations
import asyncio
import json
import sys
from typing import TYPE_CHECKING
try:
# Skip on Windows - cassandra-driver requires libev C extension which is not available
if sys.platform == "win32":
msg = "cassandra-driver not supported on Windows"
raise ImportError(msg)
from cassandra.auth import ( # pyright: ignore[reportMissingImports]
PlainTextAuthProvider,
)
from cassandra.cluster import ( # pyright: ignore[reportMissingImports]
Cluster,
)
from cassandra.concurrent import ( # pyright: ignore[reportMissingImports]
execute_concurrent_with_args,
)
CASSANDRA_AVAILABLE = True
except ImportError:
Cluster = None
PlainTextAuthProvider = None
CASSANDRA_AVAILABLE = False
from ..exceptions import BatchPipelineError
from ..logging import Logger, get_logger
from .base import _BatchPipelineMixin, log_pipeline_item, validate_table_name
if TYPE_CHECKING:
from .._types import JSONValue
from ..spiders import Spider
[docs]
class CassandraPipeline(_BatchPipelineMixin):
native_batch = True
"""
Pipeline that sends items to an Apache Cassandra database.
Args:
hosts: Cassandra cluster hosts.
keyspace: Keyspace created with a single-node replication strategy when
absent.
table: Table storing UUID, spider, JSON data, and creation timestamp.
username: Optional authentication username.
password: Optional authentication password.
port: Cassandra native protocol port.
Example::
from silkworm.pipelines import CassandraPipeline
pipeline = CassandraPipeline(
hosts=["127.0.0.1"],
keyspace="scraping",
table="items",
username="cassandra",
password="cassandra",
)
"""
[docs]
def __init__(
self,
hosts: list[str] | None = None,
keyspace: str = "scraping",
*,
table: str = "items",
username: str | None = None,
password: str | None = None,
port: int = 9042,
) -> None:
"""
Initialize CassandraPipeline.
Args:
hosts: List of Cassandra cluster hosts (default: ["127.0.0.1"])
keyspace: Keyspace name
table: Table name (default: "items")
username: Optional username for authentication
password: Optional password for authentication
port: Cassandra port (default: 9042)
"""
if not CASSANDRA_AVAILABLE:
raise ImportError(
"cassandra-driver is required for CassandraPipeline. "
"Install it with: pip install silkworm-rs[cassandra]",
)
self.hosts: list[str] = hosts or ["127.0.0.1"]
self.keyspace: str = validate_table_name(keyspace)
self.table: str = validate_table_name(table)
self.username = username
self.password = password
self.port = port
self._cluster = None
self._session = None
self._insert_statement = None
self.logger: Logger = get_logger(component="CassandraPipeline")
[docs]
async def open(self, spider: Spider) -> None:
"""Connect and create the keyspace and JSON-document table if absent."""
# Setup authentication if credentials provided
auth_provider = None
if self.username and self.password:
auth_provider = PlainTextAuthProvider( # type: ignore[misc]
username=self.username,
password=self.password,
)
# Connect to Cassandra cluster
cluster = Cluster( # type: ignore[misc]
self.hosts,
port=self.port,
auth_provider=auth_provider,
)
self._cluster = cluster
try:
session = cluster.connect()
self._session = session
session.execute(
f"""
CREATE KEYSPACE IF NOT EXISTS {self.keyspace}
WITH replication = {{'class': 'SimpleStrategy', 'replication_factor': 1}}
""",
)
session.set_keyspace(self.keyspace)
session.execute(
f"""
CREATE TABLE IF NOT EXISTS {self.table} (
id uuid PRIMARY KEY,
spider text,
data text,
created_at timestamp
)
""",
)
self._insert_statement = session.prepare(
f"INSERT INTO {self.keyspace}.{self.table} "
"(id, spider, data, created_at) VALUES (?, ?, ?, ?)"
)
except BaseException as exc:
self._cluster = None
self._session = None
try:
cluster.shutdown()
except BaseException as cleanup_exc: # noqa: BLE001
exc.add_note(f"Cassandra rollback failed: {cleanup_exc}")
raise
self.logger.info(
"Opened Cassandra pipeline",
hosts=self.hosts,
keyspace=self.keyspace,
table=self.table,
)
[docs]
async def close(self, spider: Spider) -> None:
"""Shut down the Cassandra cluster connection."""
cluster = self._cluster
self._cluster = None
self._session = None
self._insert_statement = None
if cluster:
cluster.shutdown()
self.logger.info("Closed Cassandra pipeline", table=self.table)
[docs]
async def process_item(self, item: JSONValue, spider: Spider) -> JSONValue:
"""Insert one item with a UUID, spider name, and timestamp."""
if not self._session:
raise RuntimeError("CassandraPipeline not opened")
import uuid
from datetime import UTC, datetime
# Insert item into Cassandra
self._session.execute(
f"""
INSERT INTO {self.table} (id, spider, data, created_at)
VALUES (%s, %s, %s, %s)
""",
(
uuid.uuid4(),
spider.name,
json.dumps(item, ensure_ascii=False),
datetime.now(UTC),
),
)
log_pipeline_item(
self,
"Inserted item in Cassandra",
table=self.table,
spider=spider.name,
)
return item
[docs]
async def process_items(
self, items: list[JSONValue], spider: Spider
) -> list[JSONValue]:
if not items:
return items
if not self._session or self._insert_statement is None:
raise RuntimeError("CassandraPipeline not opened")
import uuid
from datetime import UTC, datetime
arguments = [
(
uuid.uuid4(),
spider.name,
json.dumps(item, ensure_ascii=False),
datetime.now(UTC),
)
for item in items
]
results = await asyncio.to_thread(
execute_concurrent_with_args, # pyright: ignore[reportPossiblyUnboundVariable]
self._session,
self._insert_statement,
arguments,
raise_on_first_error=False,
)
failures = [result for success, result in results if not success]
if failures:
raise BatchPipelineError(
"CassandraPipeline", total=len(items), failed=len(failures)
) from failures[0]
log_pipeline_item(
self,
"Inserted item batch in Cassandra",
table=self.table,
spider=spider.name,
item_count=len(items),
)
return items