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]
placeholders = ["%s"]
# Check if we need a synthetic ID # Check if we need a synthetic ID
has_primary_key = any(field.primary for field in schema.fields) has_primary_key = any(field.primary for field in schema.fields)
if not has_primary_key: if not has_primary_key:
import uuid import uuid
columns.append("synthetic_id") columns.append("synthetic_id")
values.append(uuid.uuid4()) values.append(uuid.uuid4())
placeholders.append("%s") placeholders.append("%s")
# Process fields # Process fields for this object
for field in schema.fields: skip_object = False
safe_field_name = self.sanitize_name(field.name) for field in schema.fields:
raw_value = obj.values.get(field.name) safe_field_name = self.sanitize_name(field.name)
raw_value = value_map.get(field.name)
# Handle required fields # Handle required fields
if field.required and raw_value is None: if field.required and raw_value is None:
logger.warning(f"Required field {field.name} is missing in object") logger.warning(f"Required field {field.name} is missing in object {obj_index}")
# Continue anyway - Cassandra doesn't enforce NOT NULL # Continue anyway - Cassandra doesn't enforce NOT NULL
# Check if primary key field is NULL # Check if primary key field is NULL
if field.primary and raw_value is None: if field.primary and raw_value is None:
logger.error(f"Primary key field {field.name} cannot be NULL - skipping object") logger.error(f"Primary key field {field.name} cannot be NULL - skipping object {obj_index}")
return skip_object = True
break
# Convert value to appropriate type # Convert value to appropriate type
converted_value = self.convert_value(raw_value, field.type) converted_value = self.convert_value(raw_value, field.type)
columns.append(safe_field_name) columns.append(safe_field_name)
values.append(converted_value) values.append(converted_value)
placeholders.append("%s") placeholders.append("%s")
# Build and execute insert query # Skip this object if primary key validation failed
insert_cql = f""" if skip_object:
INSERT INTO {safe_keyspace}.{safe_table} ({', '.join(columns)}) continue
VALUES ({', '.join(placeholders)})
"""
# Debug: Show data being inserted # Build and execute insert query for this object
logger.debug(f"Storing {obj.schema_name}: {dict(zip(columns, values))}") insert_cql = f"""
INSERT INTO {safe_keyspace}.{safe_table} ({', '.join(columns)})
VALUES ({', '.join(placeholders)})
"""
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} object {obj_index}: {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: {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 {obj_index}: {e}", exc_info=True)
raise
def close(self): def close(self):
"""Clean up Cassandra connections""" """Clean up Cassandra connections"""