mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-24 20:51:02 +02:00
Object batching
This commit is contained in:
parent
ebca467ed8
commit
162859f1d6
4 changed files with 78 additions and 61 deletions
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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", ""),
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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"""
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue