Object batching

This commit is contained in:
Cyber MacGeddon 2025-09-05 15:48:04 +01:00
parent ebca467ed8
commit 162859f1d6
4 changed files with 78 additions and 61 deletions

View file

@ -1,4 +1,4 @@
from pulsar.schema import Record, String, Map, Double from pulsar.schema import Record, String, Map, Double, Array
from ..core.metadata import Metadata from ..core.metadata import Metadata
from ..core.topic import topic from ..core.topic import topic
@ -10,7 +10,7 @@ from ..core.topic import topic
class ExtractedObject(Record): class ExtractedObject(Record):
metadata = Metadata() metadata = Metadata()
schema_name = String() # Which schema this object belongs to schema_name = String() # Which schema this object belongs to
values = Map(String()) # Field name -> value values = Array(Map(String())) # Array of objects, each object is field name -> value
confidence = Double() confidence = Double()
source_span = String() # Text span where object was found source_span = String() # Text span where object was found

View file

@ -256,31 +256,34 @@ class Processor(FlowProcessor):
flow flow
) )
# Emit each extracted object # Emit extracted objects as a batch if any were found
for obj in objects: if objects:
# Calculate confidence (could be enhanced with actual confidence from prompt) # Calculate confidence (could be enhanced with actual confidence from prompt)
confidence = 0.8 # Default confidence confidence = 0.8 # Default confidence
# Convert all values to strings for Pulsar compatibility # Convert all objects' values to strings for Pulsar compatibility
string_values = convert_values_to_strings(obj) batch_values = []
for obj in objects:
string_values = convert_values_to_strings(obj)
batch_values.append(string_values)
# Create ExtractedObject # Create ExtractedObject with batched values
extracted = ExtractedObject( extracted = ExtractedObject(
metadata=Metadata( metadata=Metadata(
id=f"{v.metadata.id}:{schema_name}:{hash(str(obj))}", id=f"{v.metadata.id}:{schema_name}",
metadata=[], metadata=[],
user=v.metadata.user, user=v.metadata.user,
collection=v.metadata.collection, collection=v.metadata.collection,
), ),
schema_name=schema_name, schema_name=schema_name,
values=string_values, values=batch_values, # Array of objects
confidence=confidence, confidence=confidence,
source_span=chunk_text[:100] # First 100 chars as source reference source_span=chunk_text[:100] # First 100 chars as source reference
) )
await flow("output").send(extracted) await flow("output").send(extracted)
logger.debug(f"Emitted extracted object for schema {schema_name}") logger.debug(f"Emitted batch of {len(objects)} objects for schema {schema_name}")
except Exception as e: except Exception as e:
logger.error(f"Object extraction exception: {e}", exc_info=True) logger.error(f"Object extraction exception: {e}", exc_info=True)

View file

@ -44,6 +44,12 @@ class ObjectsImport:
data = msg.json() data = msg.json()
# Handle both single object and array of objects for backward compatibility
values_data = data["values"]
if not isinstance(values_data, list):
# Single object - wrap in array
values_data = [values_data]
elt = ExtractedObject( elt = ExtractedObject(
metadata=Metadata( metadata=Metadata(
id=data["metadata"]["id"], id=data["metadata"]["id"],
@ -52,7 +58,7 @@ class ObjectsImport:
collection=data["metadata"]["collection"], collection=data["metadata"]["collection"],
), ),
schema_name=data["schema_name"], schema_name=data["schema_name"],
values=data["values"], values=values_data,
confidence=data.get("confidence", 1.0), confidence=data.get("confidence", 1.0),
source_span=data.get("source_span", ""), source_span=data.get("source_span", ""),
) )

View file

@ -311,7 +311,7 @@ class Processor(FlowProcessor):
"""Process incoming ExtractedObject and store in Cassandra""" """Process incoming ExtractedObject and store in Cassandra"""
obj = msg.value() obj = msg.value()
logger.info(f"Storing object for schema {obj.schema_name} from {obj.metadata.id}") logger.info(f"Storing {len(obj.values)} objects for schema {obj.schema_name} from {obj.metadata.id}")
# Get schema definition # Get schema definition
schema = self.schemas.get(obj.schema_name) schema = self.schemas.get(obj.schema_name)
@ -328,59 +328,67 @@ class Processor(FlowProcessor):
safe_keyspace = self.sanitize_name(keyspace) safe_keyspace = self.sanitize_name(keyspace)
safe_table = self.sanitize_table(table_name) safe_table = self.sanitize_table(table_name)
# Build column names and values # Process each object in the batch
columns = ["collection"] for obj_index, value_map in enumerate(obj.values):
values = [obj.metadata.collection] # Build column names and values for this object
placeholders = ["%s"] columns = ["collection"]
values = [obj.metadata.collection]
# Check if we need a synthetic ID placeholders = ["%s"]
has_primary_key = any(field.primary for field in schema.fields)
if not has_primary_key:
import uuid
columns.append("synthetic_id")
values.append(uuid.uuid4())
placeholders.append("%s")
# Process fields
for field in schema.fields:
safe_field_name = self.sanitize_name(field.name)
raw_value = obj.values.get(field.name)
# Handle required fields # Check if we need a synthetic ID
if field.required and raw_value is None: has_primary_key = any(field.primary for field in schema.fields)
logger.warning(f"Required field {field.name} is missing in object") if not has_primary_key:
# Continue anyway - Cassandra doesn't enforce NOT NULL import uuid
columns.append("synthetic_id")
values.append(uuid.uuid4())
placeholders.append("%s")
# Check if primary key field is NULL # Process fields for this object
if field.primary and raw_value is None: skip_object = False
logger.error(f"Primary key field {field.name} cannot be NULL - skipping object") for field in schema.fields:
return safe_field_name = self.sanitize_name(field.name)
raw_value = value_map.get(field.name)
# Handle required fields
if field.required and raw_value is None:
logger.warning(f"Required field {field.name} is missing in object {obj_index}")
# Continue anyway - Cassandra doesn't enforce NOT NULL
# Check if primary key field is NULL
if field.primary and raw_value is None:
logger.error(f"Primary key field {field.name} cannot be NULL - skipping object {obj_index}")
skip_object = True
break
# Convert value to appropriate type
converted_value = self.convert_value(raw_value, field.type)
columns.append(safe_field_name)
values.append(converted_value)
placeholders.append("%s")
# Convert value to appropriate type # Skip this object if primary key validation failed
converted_value = self.convert_value(raw_value, field.type) if skip_object:
continue
columns.append(safe_field_name) # Build and execute insert query for this object
values.append(converted_value) insert_cql = f"""
placeholders.append("%s") INSERT INTO {safe_keyspace}.{safe_table} ({', '.join(columns)})
VALUES ({', '.join(placeholders)})
# Build and execute insert query """
insert_cql = f"""
INSERT INTO {safe_keyspace}.{safe_table} ({', '.join(columns)}) # Debug: Show data being inserted
VALUES ({', '.join(placeholders)}) logger.debug(f"Storing {obj.schema_name} object {obj_index}: {dict(zip(columns, values))}")
"""
if len(columns) != len(values) or len(columns) != len(placeholders):
# Debug: Show data being inserted raise ValueError(f"Mismatch in counts - columns: {len(columns)}, values: {len(values)}, placeholders: {len(placeholders)}")
logger.debug(f"Storing {obj.schema_name}: {dict(zip(columns, values))}")
try:
if len(columns) != len(values) or len(columns) != len(placeholders): # Convert to tuple - Cassandra driver requires tuple for parameters
raise ValueError(f"Mismatch in counts - columns: {len(columns)}, values: {len(values)}, placeholders: {len(placeholders)}") self.session.execute(insert_cql, tuple(values))
except Exception as e:
try: logger.error(f"Failed to insert object {obj_index}: {e}", exc_info=True)
# Convert to tuple - Cassandra driver requires tuple for parameters raise
self.session.execute(insert_cql, tuple(values))
except Exception as e:
logger.error(f"Failed to insert object: {e}", exc_info=True)
raise
def close(self): def close(self):
"""Clean up Cassandra connections""" """Clean up Cassandra connections"""