diff --git a/geaflow-ai/pom.xml b/geaflow-ai/pom.xml index 006215b74..9d450f890 100644 --- a/geaflow-ai/pom.xml +++ b/geaflow-ai/pom.xml @@ -121,8 +121,31 @@ geaflow-api ${project.version} + + org.apache.geaflow + geaflow-dsl-common + ${project.version} + + + org.apache.geaflow + geaflow-pipeline + ${project.version} + test + + + org.apache.geaflow + geaflow-on-local + ${project.version} + test + + + org.apache.geaflow + geaflow-dsl-runtime + ${project.version} + test + org.junit.jupiter junit-jupiter diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/adapter/MemoryGraphAdapter.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/adapter/MemoryGraphAdapter.java new file mode 100644 index 000000000..8edb2fb0f --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/adapter/MemoryGraphAdapter.java @@ -0,0 +1,1247 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.adapter; + +import java.time.Instant; +import java.time.format.DateTimeParseException; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.Comparator; +import java.util.HashMap; +import java.util.HashSet; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Set; +import java.util.TreeMap; +import org.apache.geaflow.ai.graph.io.Edge; +import org.apache.geaflow.ai.graph.io.EdgeGroup; +import org.apache.geaflow.ai.graph.io.EdgeSchema; +import org.apache.geaflow.ai.graph.io.EntityGroup; +import org.apache.geaflow.ai.graph.io.GraphSchema; +import org.apache.geaflow.ai.graph.io.MemoryGraph; +import org.apache.geaflow.ai.graph.io.Vertex; +import org.apache.geaflow.ai.graph.io.VertexGroup; +import org.apache.geaflow.ai.graph.io.VertexSchema; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.semantics.CanonicalSnapshot; +import org.apache.geaflow.ai.temporal.semantics.EventNormalizer; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; + +/** + * Projects canonical temporal snapshots to the existing in-memory graph model. + */ +public final class MemoryGraphAdapter { + + private static final String ENTITY = "entity"; + private static final String FACT_VERSION = "fact_version"; + private static final String MEMORY_EVENT = "memory_event"; + private static final String EVIDENCE = "evidence"; + private static final String SOURCE = "source"; + + private static final String SUBJECT = "subject"; + private static final String OBJECT = "object"; + private static final String GENERATES = "generates"; + private static final String SUPPORTED_BY = "supported_by"; + private static final String FROM_SOURCE = "from_source"; + private static final String SUPERSEDES = "supersedes"; + private static final String DUPLICATE_OF = "duplicate_of"; + private static final String CONFLICTS_WITH = "conflicts_with"; + + private static final String ENTITY_PREFIX = "entity:"; + private static final String VERSION_PREFIX = "version:"; + private static final String EVENT_PREFIX = "event:"; + private static final String EVIDENCE_PREFIX = "evidence:"; + private static final String SOURCE_PREFIX = "source:"; + private static final String EMPTY = ""; + + private static final List ENTITY_FIELDS = fields("label"); + private static final List VERSION_FIELDS = fields( + "factId", + "predicate", + "scope", + "valueKind", + "literalValue", + "status", + "validStart", + "validEnd", + "transactionStart", + "transactionEnd"); + private static final List EVENT_FIELDS = fields( + "operation", + "factId", + "subjectId", + "predicate", + "scope", + "valueKind", + "value", + "validStart", + "validEnd", + "recordedAt", + "payloadHash"); + private static final List EVIDENCE_FIELDS = fields("content"); + private static final List SOURCE_FIELDS = fields("name"); + private static final List NO_FIELDS = Collections.emptyList(); + private static final List COUNT_FIELDS = fields( + "occurrenceCount"); + private static final List RELATION_FIELDS = fields( + "relationId", "occurrenceCount"); + + private static final List VERTEX_LABELS = fields( + ENTITY, FACT_VERSION, MEMORY_EVENT, EVIDENCE, SOURCE); + private static final List EDGE_LABELS = fields( + SUBJECT, + OBJECT, + GENERATES, + SUPPORTED_BY, + FROM_SOURCE, + SUPERSEDES, + DUPLICATE_OF, + CONFLICTS_WITH); + + private static final Comparator VERTEX_ORDER = + Comparator.comparing(Vertex::getId); + private static final Comparator EDGE_ORDER = + Comparator.comparing(Edge::getSrcId) + .thenComparing(Edge::getDstId) + .thenComparing(edge -> edge.getValues().toString()); + + public MemoryGraph toGraph(CanonicalSnapshot snapshot) { + Objects.requireNonNull(snapshot, "snapshot"); + + Map memoryEntities = new TreeMap<>(); + Map evidenceById = new TreeMap<>(); + Map sources = new TreeMap<>(); + Map versionVertices = new TreeMap<>(); + Map eventVertices = new TreeMap<>(); + Map> edges = emptyEdgeLists(); + Map supported = new TreeMap<>(); + Map> relationEdges = + relationEdgeMaps(); + + for (NormalizedMemoryEvent event : snapshot.getEvents()) { + collectEventEntities(event, memoryEntities); + collectEvidence(event.getEvidence(), evidenceById, sources); + String eventGraphId = eventId(event.getEventId()); + putUniqueVertex( + eventVertices, + new Vertex(MEMORY_EVENT, eventGraphId, eventValues(event))); + for (Evidence evidence : event.getEvidence()) { + addCountedEdge( + supported, + eventGraphId, + evidenceId(evidence.getId())); + } + } + + for (Map.Entry> entry + : snapshot.getState().getVersionsByFactKey().entrySet()) { + FactKey key = entry.getKey(); + for (MemoryFactVersion version : entry.getValue()) { + MemoryFact fact = version.getFact(); + collectEntity(memoryEntities, fact.getSubject()); + if (fact.isRelationship()) { + collectEntity(memoryEntities, fact.getTarget().get()); + } + collectEvidence(version.getEvidence(), evidenceById, sources); + + String versionGraphId = versionId(version.getId()); + putUniqueVertex( + versionVertices, + new Vertex( + FACT_VERSION, + versionGraphId, + versionValues(version, key))); + edges.get(SUBJECT).add(new Edge( + SUBJECT, + versionGraphId, + entityId(fact.getSubject().getId()), + NO_FIELDS)); + if (fact.isRelationship()) { + edges.get(OBJECT).add(new Edge( + OBJECT, + versionGraphId, + entityId(fact.getTarget().get().getId()), + NO_FIELDS)); + } + String generatingEventId = + snapshot.getGeneratingEventIds().get(version.getId()); + edges.get(GENERATES).add(new Edge( + GENERATES, + eventId(generatingEventId), + versionGraphId, + NO_FIELDS)); + for (Evidence evidence : version.getEvidence()) { + addCountedEdge( + supported, + versionGraphId, + evidenceId(evidence.getId())); + } + } + } + + for (VersionRelation relation + : snapshot.getState().getRelations()) { + String label = relationLabel(relation.getType()); + addCountedEdge( + relationEdges.get(label), + versionId(relation.getFromVersionId()), + versionId(relation.getToVersionId())); + } + + List entityVertices = new ArrayList<>(); + for (MemoryEntity entity : memoryEntities.values()) { + entityVertices.add(new Vertex( + ENTITY, + entityId(entity.getId()), + fields(entity.getLabel()))); + } + List evidenceVertices = new ArrayList<>(); + for (Evidence evidence : evidenceById.values()) { + evidenceVertices.add(new Vertex( + EVIDENCE, + evidenceId(evidence.getId()), + fields(evidence.getContent()))); + edges.get(FROM_SOURCE).add(new Edge( + FROM_SOURCE, + evidenceId(evidence.getId()), + sourceId(evidence.getSource().getId()), + NO_FIELDS)); + } + List sourceVertices = new ArrayList<>(); + for (Source source : sources.values()) { + sourceVertices.add(new Vertex( + SOURCE, + sourceId(source.getId()), + fields(source.getName()))); + } + + edges.put( + SUPPORTED_BY, + countedEdges(SUPPORTED_BY, supported, false)); + for (String label : Arrays.asList( + SUPERSEDES, DUPLICATE_OF, CONFLICTS_WITH)) { + edges.put( + label, + countedEdges(label, relationEdges.get(label), true)); + } + + Map> vertices = new LinkedHashMap<>(); + vertices.put(ENTITY, entityVertices); + vertices.put( + FACT_VERSION, + new ArrayList<>(versionVertices.values())); + vertices.put( + MEMORY_EVENT, + new ArrayList<>(eventVertices.values())); + vertices.put(EVIDENCE, evidenceVertices); + vertices.put(SOURCE, sourceVertices); + return createGraph(vertices, edges); + } + + public CanonicalSnapshot fromGraph(MemoryGraph graph) { + Objects.requireNonNull(graph, "graph"); + validateGraphShape(graph); + + Map entityVertices = + vertexIndex(graph, ENTITY, ENTITY_PREFIX, ENTITY_FIELDS); + Map versionVertices = + vertexIndex( + graph, + FACT_VERSION, + VERSION_PREFIX, + VERSION_FIELDS); + Map eventVertices = + vertexIndex( + graph, + MEMORY_EVENT, + EVENT_PREFIX, + EVENT_FIELDS); + Map evidenceVertices = + vertexIndex( + graph, + EVIDENCE, + EVIDENCE_PREFIX, + EVIDENCE_FIELDS); + Map sourceVertices = + vertexIndex(graph, SOURCE, SOURCE_PREFIX, SOURCE_FIELDS); + Map> edges = edgeIndex(graph); + + Map memoryEntities = + readEntities(entityVertices); + Map sources = readSources(sourceVertices); + Map evidence = readEvidence( + evidenceVertices, + sources, + edges.get(FROM_SOURCE)); + Map> evidenceByOwner = + readSupportedEvidence( + edges.get(SUPPORTED_BY), + eventVertices, + versionVertices, + evidence); + + Map events = readEvents( + eventVertices, + memoryEntities, + evidenceByOwner); + Map> versions = readVersions( + versionVertices, + memoryEntities, + evidenceByOwner, + edges.get(SUBJECT), + edges.get(OBJECT)); + validateNoOrphanVertices( + memoryEntities, + evidence, + sources, + events, + versions); + Map generating = readGeneratingEvents( + edges.get(GENERATES), + events, + versionVertices); + List relations = readRelations( + edges, + versionVertices); + + return new CanonicalSnapshot( + new TemporalState(versions, relations), + new ArrayList<>(events.values()), + generating); + } + + private static MemoryGraph createGraph( + Map> vertices, + Map> edges) { + GraphSchema schema = new GraphSchema(); + Map vertexSchemas = new LinkedHashMap<>(); + addVertexSchema(schema, vertexSchemas, ENTITY, ENTITY_FIELDS); + addVertexSchema( + schema, vertexSchemas, FACT_VERSION, VERSION_FIELDS); + addVertexSchema( + schema, vertexSchemas, MEMORY_EVENT, EVENT_FIELDS); + addVertexSchema(schema, vertexSchemas, EVIDENCE, EVIDENCE_FIELDS); + addVertexSchema(schema, vertexSchemas, SOURCE, SOURCE_FIELDS); + + Map edgeSchemas = new LinkedHashMap<>(); + addEdgeSchema(schema, edgeSchemas, SUBJECT, NO_FIELDS); + addEdgeSchema(schema, edgeSchemas, OBJECT, NO_FIELDS); + addEdgeSchema(schema, edgeSchemas, GENERATES, NO_FIELDS); + addEdgeSchema(schema, edgeSchemas, SUPPORTED_BY, COUNT_FIELDS); + addEdgeSchema(schema, edgeSchemas, FROM_SOURCE, NO_FIELDS); + addEdgeSchema(schema, edgeSchemas, SUPERSEDES, RELATION_FIELDS); + addEdgeSchema(schema, edgeSchemas, DUPLICATE_OF, RELATION_FIELDS); + addEdgeSchema( + schema, edgeSchemas, CONFLICTS_WITH, RELATION_FIELDS); + + Map groups = new LinkedHashMap<>(); + for (String label : VERTEX_LABELS) { + List rows = vertices.get(label); + Collections.sort(rows, VERTEX_ORDER); + groups.put( + label, + new VertexGroup(vertexSchemas.get(label), rows)); + } + for (String label : EDGE_LABELS) { + List rows = edges.get(label); + Collections.sort(rows, EDGE_ORDER); + groups.put( + label, + new EdgeGroup(edgeSchemas.get(label), rows)); + } + return new MemoryGraph(schema, groups); + } + + private static void addVertexSchema( + GraphSchema graphSchema, + Map schemas, + String label, + List fields) { + VertexSchema schema = new VertexSchema(label, "id", fields); + graphSchema.addVertex(schema); + schemas.put(label, schema); + } + + private static void addEdgeSchema( + GraphSchema graphSchema, + Map schemas, + String label, + List fields) { + EdgeSchema schema = new EdgeSchema( + label, "srcId", "dstId", fields); + graphSchema.addEdge(schema); + schemas.put(label, schema); + } + + private static Map> emptyEdgeLists() { + Map> edges = new LinkedHashMap<>(); + for (String label : EDGE_LABELS) { + edges.put(label, new ArrayList<>()); + } + return edges; + } + + private static Map> + relationEdgeMaps() { + Map> relations = + new LinkedHashMap<>(); + relations.put(SUPERSEDES, new TreeMap<>()); + relations.put(DUPLICATE_OF, new TreeMap<>()); + relations.put(CONFLICTS_WITH, new TreeMap<>()); + return relations; + } + + private static List countedEdges( + String label, + Map counted, + boolean relation) { + List edges = new ArrayList<>(); + for (CountedEdge item : counted.values()) { + List values = relation + ? fields( + relationId(label, item.sourceId, item.targetId), + Integer.toString(item.count)) + : fields(Integer.toString(item.count)); + edges.add(new Edge( + label, item.sourceId, item.targetId, values)); + } + return edges; + } + + private static void addCountedEdge( + Map edges, + String sourceId, + String targetId) { + String key = tuple(sourceId, targetId); + CountedEdge current = edges.get(key); + if (current == null) { + edges.put(key, new CountedEdge(sourceId, targetId)); + } else { + current.count++; + } + } + + private static List eventValues(NormalizedMemoryEvent event) { + String kind = EMPTY; + String value = EMPTY; + if (event.getFactValue().isPresent()) { + kind = event.getFactValue().get().getKind().name(); + value = event.getFactValue().get().getValue(); + } + return fields( + event.getOperation().name(), + event.getFactId(), + event.getFactKey().getSubjectId(), + event.getFactKey().getPredicate(), + event.getFactKey().getScope(), + kind, + value, + instant(event.getValidTime().getStart()), + optionalInstant(event.getValidTime()), + instant(event.getRecordedAt()), + event.getPayloadHash()); + } + + private static List versionValues( + MemoryFactVersion version, + FactKey key) { + MemoryFact fact = version.getFact(); + String kind = fact.isRelationship() + ? FactValue.Kind.ENTITY_REF.name() + : FactValue.Kind.LITERAL.name(); + String literalValue = fact.getLiteralValue().orElse(EMPTY); + return fields( + fact.getId(), + fact.getPredicate(), + key.getScope(), + kind, + literalValue, + version.getStatus().name(), + instant(version.getValidTime().getStart()), + optionalInstant(version.getValidTime()), + instant(version.getTransactionTime().getStart()), + optionalInstant(version.getTransactionTime())); + } + + private static void collectEventEntities( + NormalizedMemoryEvent event, + Map entities) { + if (!event.getEvent().getFact().isPresent()) { + return; + } + MemoryFact fact = event.getEvent().getFact().get(); + collectEntity(entities, fact.getSubject()); + if (fact.isRelationship()) { + collectEntity(entities, fact.getTarget().get()); + } + } + + private static void collectEntity( + Map entities, + MemoryEntity entity) { + putUnique(entities, entity.getId(), entity, "entity"); + } + + private static void collectEvidence( + List evidence, + Map evidenceById, + Map sources) { + for (Evidence item : evidence) { + putUnique(evidenceById, item.getId(), item, "evidence"); + putUnique( + sources, + item.getSource().getId(), + item.getSource(), + "source"); + } + } + + private static void putUnique( + Map values, + String id, + T value, + String type) { + T previous = values.get(id); + if (previous != null && !previous.equals(value)) { + throw new IllegalArgumentException( + "Conflicting " + type + " id: " + id); + } + values.put(id, value); + } + + private static void putUniqueVertex( + Map vertices, + Vertex vertex) { + Vertex previous = vertices.put(vertex.getId(), vertex); + if (previous != null + && !previous.getValues().equals(vertex.getValues())) { + throw new IllegalArgumentException( + "Conflicting vertex id: " + vertex.getId()); + } + } + + private static void validateGraphShape(MemoryGraph graph) { + GraphSchema schema = Objects.requireNonNull( + graph.getGraphSchema(), "graphSchema"); + require( + schema.getVertexSchemaList().size() == VERTEX_LABELS.size(), + "Unexpected vertex schema count"); + require( + schema.getEdgeSchemaList().size() == EDGE_LABELS.size(), + "Unexpected edge schema count"); + + for (int index = 0; index < VERTEX_LABELS.size(); index++) { + String label = VERTEX_LABELS.get(index); + VertexSchema actual = schema.getVertexSchemaList().get(index); + require(label.equals(actual.getLabel()), + "Unexpected vertex schema: " + actual.getLabel()); + require("id".equals(actual.getIdField()), + "Unexpected vertex id field: " + label); + require(vertexFields(label).equals(actual.getFields()), + "Unexpected vertex fields: " + label); + } + for (int index = 0; index < EDGE_LABELS.size(); index++) { + String label = EDGE_LABELS.get(index); + EdgeSchema actual = schema.getEdgeSchemaList().get(index); + require(label.equals(actual.getLabel()), + "Unexpected edge schema: " + actual.getLabel()); + require("srcId".equals(actual.getSrcIdField()), + "Unexpected edge source field: " + label); + require("dstId".equals(actual.getDstIdField()), + "Unexpected edge target field: " + label); + require(edgeFields(label).equals(actual.getFields()), + "Unexpected edge fields: " + label); + } + + Map groups = Objects.requireNonNull( + graph.entities, "graph.entities"); + List expectedGroups = new ArrayList<>(VERTEX_LABELS); + expectedGroups.addAll(EDGE_LABELS); + require( + expectedGroups.equals(new ArrayList<>(groups.keySet())), + "Unexpected graph entity groups"); + for (String label : VERTEX_LABELS) { + require(groups.get(label) instanceof VertexGroup, + "Expected vertex group: " + label); + } + for (String label : EDGE_LABELS) { + require(groups.get(label) instanceof EdgeGroup, + "Expected edge group: " + label); + } + } + + private static Map vertexIndex( + MemoryGraph graph, + String label, + String prefix, + List expectedFields) { + Map result = new TreeMap<>(); + for (Vertex vertex + : ((VertexGroup) graph.entities.get(label)).getVertices()) { + require(vertex != null, "Null vertex in group: " + label); + require(label.equals(vertex.getLabel()), + "Vertex label does not match group: " + vertex.getId()); + rawId(vertex.getId(), prefix); + require(vertex.getValues() != null, + "Null vertex values: " + vertex.getId()); + require(vertex.getValues().size() == expectedFields.size(), + "Unexpected vertex value count: " + vertex.getId()); + for (String fieldValue : vertex.getValues()) { + require(fieldValue != null, + "Null vertex field: " + vertex.getId()); + } + require(result.put(vertex.getId(), vertex) == null, + "Duplicate vertex id: " + vertex.getId()); + } + return result; + } + + private static Map> edgeIndex(MemoryGraph graph) { + Map> result = new LinkedHashMap<>(); + for (String label : EDGE_LABELS) { + List rows = new ArrayList<>(); + Set identities = new HashSet<>(); + for (Edge edge + : ((EdgeGroup) graph.entities.get(label)).getOutEdges()) { + require(edge != null, "Null edge in group: " + label); + require(label.equals(edge.getLabel()), + "Edge label does not match group: " + label); + require(edge.getSrcId() != null && edge.getDstId() != null, + "Null edge endpoint: " + label); + require(edge.getValues() != null, + "Null edge values: " + label); + require(edge.getValues().size() == edgeFields(label).size(), + "Unexpected edge value count: " + label); + for (String fieldValue : edge.getValues()) { + require(fieldValue != null, + "Null edge field: " + label); + } + require(identities.add(edge), + "Duplicate edge identity: " + edge); + rows.add(edge); + } + Collections.sort(rows, EDGE_ORDER); + result.put(label, rows); + } + return result; + } + + private static Map readEntities( + Map vertices) { + Map result = new TreeMap<>(); + for (Vertex vertex : vertices.values()) { + String rawId = rawId(vertex.getId(), ENTITY_PREFIX); + result.put( + vertex.getId(), + new MemoryEntity( + rawId, + value(vertex, ENTITY_FIELDS, "label"))); + } + return result; + } + + private static Map readSources( + Map vertices) { + Map result = new TreeMap<>(); + for (Vertex vertex : vertices.values()) { + String rawId = rawId(vertex.getId(), SOURCE_PREFIX); + result.put( + vertex.getId(), + new Source( + rawId, + value(vertex, SOURCE_FIELDS, "name"))); + } + return result; + } + + private static Map readEvidence( + Map vertices, + Map sources, + List fromSourceEdges) { + Map sourceByEvidence = new HashMap<>(); + for (Edge edge : fromSourceEdges) { + require(vertices.containsKey(edge.getSrcId()), + "Unknown evidence in from_source edge"); + require(sources.containsKey(edge.getDstId()), + "Unknown source in from_source edge"); + require(sourceByEvidence.put( + edge.getSrcId(), edge.getDstId()) == null, + "Evidence has multiple sources: " + edge.getSrcId()); + } + require(sourceByEvidence.keySet().equals(vertices.keySet()), + "Every evidence must have exactly one source"); + + Map result = new TreeMap<>(); + for (Vertex vertex : vertices.values()) { + String rawId = rawId(vertex.getId(), EVIDENCE_PREFIX); + result.put( + vertex.getId(), + new Evidence( + rawId, + sources.get(sourceByEvidence.get(vertex.getId())), + value(vertex, EVIDENCE_FIELDS, "content"))); + } + return result; + } + + private static Map> readSupportedEvidence( + List supportedEdges, + Map eventVertices, + Map versionVertices, + Map evidence) { + Map> result = new HashMap<>(); + for (Edge edge : supportedEdges) { + require( + eventVertices.containsKey(edge.getSrcId()) + || versionVertices.containsKey(edge.getSrcId()), + "Unknown supported_by owner: " + edge.getSrcId()); + Evidence item = evidence.get(edge.getDstId()); + require(item != null, + "Unknown supported_by evidence: " + edge.getDstId()); + int count = positiveCount(edge.getValues().get(0)); + List ownerEvidence = result.computeIfAbsent( + edge.getSrcId(), ignored -> new ArrayList<>()); + for (int occurrence = 0; occurrence < count; occurrence++) { + ownerEvidence.add(item); + } + } + return result; + } + + private static Map readEvents( + Map vertices, + Map entities, + Map> evidenceByOwner) { + Map result = new TreeMap<>(); + EventNormalizer normalizer = new EventNormalizer(); + for (Vertex vertex : vertices.values()) { + String rawEventId = rawId(vertex.getId(), EVENT_PREFIX); + MemoryEventOperation operation = enumValue( + MemoryEventOperation.class, + value(vertex, EVENT_FIELDS, "operation"), + "event operation"); + String factId = value(vertex, EVENT_FIELDS, "factId"); + String subjectId = value(vertex, EVENT_FIELDS, "subjectId"); + String predicate = value(vertex, EVENT_FIELDS, "predicate"); + FactKey key = new FactKey( + subjectId, + predicate, + value(vertex, EVENT_FIELDS, "scope")); + TimeInterval validTime = interval( + value(vertex, EVENT_FIELDS, "validStart"), + value(vertex, EVENT_FIELDS, "validEnd")); + Instant recordedAt = parseInstant( + value(vertex, EVENT_FIELDS, "recordedAt")); + List eventEvidence = evidenceByOwner.get( + vertex.getId()); + require(eventEvidence != null && !eventEvidence.isEmpty(), + "Memory event must have evidence: " + rawEventId); + + MemoryEvent event; + if (operation == MemoryEventOperation.RETRACT) { + require(value(vertex, EVENT_FIELDS, "valueKind").isEmpty(), + "Retract event must not have a value kind"); + require(value(vertex, EVENT_FIELDS, "value").isEmpty(), + "Retract event must not have a value"); + event = MemoryEvent.retract( + rawEventId, + factId, + validTime, + recordedAt, + eventEvidence); + } else { + MemoryFact fact = eventFact( + vertex, + factId, + subjectId, + predicate, + entities); + if (operation == MemoryEventOperation.ADD) { + event = MemoryEvent.add( + rawEventId, + fact, + validTime, + recordedAt, + eventEvidence); + } else { + require(operation == MemoryEventOperation.CORRECT, + "Unsupported memory event operation"); + event = MemoryEvent.correct( + rawEventId, + fact, + validTime, + recordedAt, + eventEvidence); + } + } + + NormalizedMemoryEvent normalized = normalizer.normalize(event, key); + require(eventId(normalized.getEventId()).equals(vertex.getId()), + "Event id is not canonical: " + rawEventId); + require( + normalized.getPayloadHash().equals( + value(vertex, EVENT_FIELDS, "payloadHash")), + "Event payload hash does not match: " + rawEventId); + result.put(vertex.getId(), normalized); + } + return result; + } + + private static MemoryFact eventFact( + Vertex vertex, + String factId, + String subjectId, + String predicate, + Map entities) { + MemoryEntity subject = entities.get(entityId(subjectId)); + require(subject != null, + "Unknown event subject: " + subjectId); + FactValue.Kind kind = enumValue( + FactValue.Kind.class, + value(vertex, EVENT_FIELDS, "valueKind"), + "event value kind"); + String factValue = value(vertex, EVENT_FIELDS, "value"); + if (kind == FactValue.Kind.LITERAL) { + return MemoryFact.attribute( + factId, subject, predicate, factValue); + } + MemoryEntity target = entities.get(entityId(factValue)); + require(target != null, + "Unknown event object: " + factValue); + return MemoryFact.relationship( + factId, subject, predicate, target); + } + + private static Map> readVersions( + Map vertices, + Map entities, + Map> evidenceByOwner, + List subjectEdges, + List objectEdges) { + Map subjects = uniqueTargets( + subjectEdges, + vertices, + entities, + "subject"); + require(subjects.keySet().equals(vertices.keySet()), + "Every fact version must have exactly one subject"); + Map objects = uniqueTargets( + objectEdges, + vertices, + entities, + "object"); + + Map> result = new TreeMap<>(); + for (Vertex vertex : vertices.values()) { + String rawVersionId = rawId(vertex.getId(), VERSION_PREFIX); + MemoryEntity subject = entities.get(subjects.get(vertex.getId())); + String predicate = value(vertex, VERSION_FIELDS, "predicate"); + FactKey key = new FactKey( + subject.getId(), + predicate, + value(vertex, VERSION_FIELDS, "scope")); + + FactValue.Kind kind = enumValue( + FactValue.Kind.class, + value(vertex, VERSION_FIELDS, "valueKind"), + "version value kind"); + String factId = value(vertex, VERSION_FIELDS, "factId"); + String literalValue = value( + vertex, VERSION_FIELDS, "literalValue"); + MemoryFact fact; + if (kind == FactValue.Kind.LITERAL) { + require(!objects.containsKey(vertex.getId()), + "Literal version must not have an object"); + fact = MemoryFact.attribute( + factId, subject, predicate, literalValue); + } else { + require(literalValue.isEmpty(), + "Entity-reference version must not have a literal value"); + String objectId = objects.get(vertex.getId()); + require(objectId != null, + "Entity-reference version must have an object"); + fact = MemoryFact.relationship( + factId, subject, predicate, entities.get(objectId)); + } + + List versionEvidence = evidenceByOwner.get( + vertex.getId()); + require(versionEvidence != null && !versionEvidence.isEmpty(), + "Fact version must have evidence: " + rawVersionId); + MemoryFactVersion version = new MemoryFactVersion( + rawVersionId, + fact, + enumValue( + MemoryFactVersionStatus.class, + value(vertex, VERSION_FIELDS, "status"), + "version status"), + interval( + value(vertex, VERSION_FIELDS, "validStart"), + value(vertex, VERSION_FIELDS, "validEnd")), + interval( + value(vertex, VERSION_FIELDS, "transactionStart"), + value(vertex, VERSION_FIELDS, "transactionEnd")), + versionEvidence); + result.computeIfAbsent( + key, ignored -> new ArrayList<>()).add(version); + } + return result; + } + + private static Map uniqueTargets( + List edges, + Map sourceVertices, + Map targets, + String relationName) { + Map result = new HashMap<>(); + for (Edge edge : edges) { + require(sourceVertices.containsKey(edge.getSrcId()), + "Unknown " + relationName + " source"); + require(targets.containsKey(edge.getDstId()), + "Unknown " + relationName + " target"); + require(result.put(edge.getSrcId(), edge.getDstId()) == null, + "Multiple " + relationName + " targets"); + } + return result; + } + + private static void validateNoOrphanVertices( + Map entities, + Map evidence, + Map sources, + Map events, + Map> versions) { + Map referencedEntities = new TreeMap<>(); + Map referencedEvidence = new TreeMap<>(); + Map referencedSources = new TreeMap<>(); + for (NormalizedMemoryEvent event : events.values()) { + collectEventEntities(event, referencedEntities); + collectEvidence( + event.getEvidence(), + referencedEvidence, + referencedSources); + } + for (List factVersions : versions.values()) { + for (MemoryFactVersion version : factVersions) { + MemoryFact fact = version.getFact(); + collectEntity(referencedEntities, fact.getSubject()); + if (fact.isRelationship()) { + collectEntity( + referencedEntities, + fact.getTarget().get()); + } + collectEvidence( + version.getEvidence(), + referencedEvidence, + referencedSources); + } + } + requireAllReferenced( + entities, referencedEntities, ENTITY_PREFIX, "entity"); + requireAllReferenced( + evidence, referencedEvidence, EVIDENCE_PREFIX, "evidence"); + requireAllReferenced( + sources, referencedSources, SOURCE_PREFIX, "source"); + } + + private static void requireAllReferenced( + Map graphValues, + Map referencedValues, + String prefix, + String type) { + for (String graphId : graphValues.keySet()) { + require(referencedValues.containsKey(rawId(graphId, prefix)), + "Orphan " + type + " vertex: " + graphId); + } + } + + private static Map readGeneratingEvents( + List generateEdges, + Map events, + Map versions) { + Map byVersion = new HashMap<>(); + for (Edge edge : generateEdges) { + require(events.containsKey(edge.getSrcId()), + "Unknown generating event: " + edge.getSrcId()); + require(versions.containsKey(edge.getDstId()), + "Unknown generated version: " + edge.getDstId()); + require(byVersion.put( + edge.getDstId(), edge.getSrcId()) == null, + "Version has multiple generating events: " + + edge.getDstId()); + } + require(byVersion.keySet().equals(versions.keySet()), + "Every fact version must have one generating event"); + + Map result = new TreeMap<>(); + for (Map.Entry entry : byVersion.entrySet()) { + result.put( + rawId(entry.getKey(), VERSION_PREFIX), + rawId(entry.getValue(), EVENT_PREFIX)); + } + return result; + } + + private static List readRelations( + Map> edges, + Map versions) { + List result = new ArrayList<>(); + for (String label : Arrays.asList( + SUPERSEDES, DUPLICATE_OF, CONFLICTS_WITH)) { + for (Edge edge : edges.get(label)) { + require(versions.containsKey(edge.getSrcId()), + "Unknown relation source: " + edge.getSrcId()); + require(versions.containsKey(edge.getDstId()), + "Unknown relation target: " + edge.getDstId()); + require( + relationId(label, edge.getSrcId(), edge.getDstId()) + .equals(edge.getValues().get(0)), + "Version relation id does not match endpoints"); + int count = positiveCount(edge.getValues().get(1)); + for (int occurrence = 0; occurrence < count; occurrence++) { + VersionRelation relation = new VersionRelation( + relationType(label), + rawId(edge.getSrcId(), VERSION_PREFIX), + rawId(edge.getDstId(), VERSION_PREFIX)); + require( + versionId(relation.getFromVersionId()).equals( + edge.getSrcId()) + && versionId(relation.getToVersionId()).equals( + edge.getDstId()), + "Version relation endpoints are not canonical"); + result.add(relation); + } + } + } + return result; + } + + private static List vertexFields(String label) { + if (ENTITY.equals(label)) { + return ENTITY_FIELDS; + } + if (FACT_VERSION.equals(label)) { + return VERSION_FIELDS; + } + if (MEMORY_EVENT.equals(label)) { + return EVENT_FIELDS; + } + if (EVIDENCE.equals(label)) { + return EVIDENCE_FIELDS; + } + if (SOURCE.equals(label)) { + return SOURCE_FIELDS; + } + throw new IllegalArgumentException( + "Unknown vertex label: " + label); + } + + private static List edgeFields(String label) { + if (SUPPORTED_BY.equals(label)) { + return COUNT_FIELDS; + } + if (SUPERSEDES.equals(label) + || DUPLICATE_OF.equals(label) + || CONFLICTS_WITH.equals(label)) { + return RELATION_FIELDS; + } + if (SUBJECT.equals(label) + || OBJECT.equals(label) + || GENERATES.equals(label) + || FROM_SOURCE.equals(label)) { + return NO_FIELDS; + } + throw new IllegalArgumentException( + "Unknown edge label: " + label); + } + + private static String value( + Vertex vertex, + List fields, + String field) { + int index = fields.indexOf(field); + if (index < 0) { + throw new IllegalArgumentException("Unknown field: " + field); + } + return vertex.getValues().get(index); + } + + private static TimeInterval interval(String start, String end) { + Instant parsedStart = parseInstant(start); + return end.isEmpty() + ? TimeInterval.unboundedFrom(parsedStart) + : new TimeInterval(parsedStart, parseInstant(end)); + } + + private static Instant parseInstant(String value) { + try { + return Instant.parse(value); + } catch (DateTimeParseException exception) { + throw new IllegalArgumentException( + "Invalid instant: " + value, + exception); + } + } + + private static int positiveCount(String value) { + try { + int count = Integer.parseInt(value); + require(count > 0, "Occurrence count must be positive"); + return count; + } catch (NumberFormatException exception) { + throw new IllegalArgumentException( + "Invalid occurrence count: " + value, + exception); + } + } + + private static > T enumValue( + Class type, + String value, + String fieldName) { + try { + return Enum.valueOf(type, value); + } catch (IllegalArgumentException exception) { + throw new IllegalArgumentException( + "Invalid " + fieldName + ": " + value, + exception); + } + } + + private static String relationLabel(VersionRelationType type) { + if (type == VersionRelationType.SUPERSEDES) { + return SUPERSEDES; + } + if (type == VersionRelationType.DUPLICATE_OF) { + return DUPLICATE_OF; + } + if (type == VersionRelationType.CONFLICTS_WITH) { + return CONFLICTS_WITH; + } + throw new IllegalArgumentException( + "Unsupported version relation type: " + type); + } + + private static VersionRelationType relationType(String label) { + if (SUPERSEDES.equals(label)) { + return VersionRelationType.SUPERSEDES; + } + if (DUPLICATE_OF.equals(label)) { + return VersionRelationType.DUPLICATE_OF; + } + if (CONFLICTS_WITH.equals(label)) { + return VersionRelationType.CONFLICTS_WITH; + } + throw new IllegalArgumentException( + "Unsupported version relation label: " + label); + } + + private static String relationId( + String label, + String sourceId, + String targetId) { + return "relation:" + tuple(label, sourceId, targetId); + } + + private static String tuple(String... values) { + StringBuilder result = new StringBuilder(); + for (String value : values) { + result.append(value.length()).append(':').append(value); + } + return result.toString(); + } + + private static String instant(Instant instant) { + return instant.toString(); + } + + private static String optionalInstant(TimeInterval interval) { + return interval.getEnd().isPresent() + ? instant(interval.getEnd().get()) : EMPTY; + } + + private static String entityId(String rawId) { + return ENTITY_PREFIX + rawId; + } + + private static String versionId(String rawId) { + return VERSION_PREFIX + rawId; + } + + private static String eventId(String rawId) { + return EVENT_PREFIX + rawId; + } + + private static String evidenceId(String rawId) { + return EVIDENCE_PREFIX + rawId; + } + + private static String sourceId(String rawId) { + return SOURCE_PREFIX + rawId; + } + + private static String rawId(String graphId, String prefix) { + require(graphId != null && graphId.startsWith(prefix), + "Graph id has the wrong prefix: " + graphId); + String rawId = graphId.substring(prefix.length()); + require(!rawId.trim().isEmpty(), + "Graph id has an empty raw id: " + graphId); + return rawId; + } + + private static List fields(String... values) { + return Collections.unmodifiableList(Arrays.asList(values)); + } + + private static void require(boolean condition, String message) { + if (!condition) { + throw new IllegalArgumentException(message); + } + } + + private static final class CountedEdge { + + private final String sourceId; + private final String targetId; + private int count; + + private CountedEdge(String sourceId, String targetId) { + this.sourceId = sourceId; + this.targetId = targetId; + this.count = 1; + } + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/baseline/LwwBaseline.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/baseline/LwwBaseline.java new file mode 100644 index 000000000..cdcfd7ef3 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/baseline/LwwBaseline.java @@ -0,0 +1,140 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.baseline; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.TreeMap; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.oracle.ReplayMethod; +import org.apache.geaflow.ai.temporal.semantics.CanonicalSnapshot; +import org.apache.geaflow.ai.temporal.semantics.EventLedger; +import org.apache.geaflow.ai.temporal.semantics.EventLedgerDecision; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; + +/** + * Replays events using deterministic last-write-wins current-state semantics. + */ +public final class LwwBaseline implements ReplayMethod { + + private static final Comparator EVENT_ORDER = + Comparator.comparing(NormalizedMemoryEvent::getRecordedAt) + .thenComparing(NormalizedMemoryEvent::getEventId); + + private static final TimeInterval ALL_VALID_TIME = + TimeInterval.unboundedFrom( + Instant.ofEpochMilli(Long.MIN_VALUE)); + + @Override + public CanonicalSnapshot replayToSnapshot( + List events) { + List ordered = new ArrayList<>( + Objects.requireNonNull(events, "events")); + for (NormalizedMemoryEvent event : ordered) { + Objects.requireNonNull(event, "event"); + } + Collections.sort(ordered, EVENT_ORDER); + + EventLedger ledger = new EventLedger(); + List accepted = new ArrayList<>(); + Map current = new TreeMap<>(); + Map generatingEventIds = new TreeMap<>(); + for (NormalizedMemoryEvent event : ordered) { + EventLedgerDecision decision = ledger.check(event); + if (decision == EventLedgerDecision.DUPLICATE_NOOP) { + continue; + } + if (decision == EventLedgerDecision.REJECT_EVENT_ID_REUSE) { + throw new IllegalArgumentException( + "Event id reused with a different payload: " + + event.getEventId()); + } + + apply(event, current, generatingEventIds); + ledger.commit(event); + accepted.add(event); + } + + Map> versions = + new TreeMap<>(); + for (Map.Entry entry + : current.entrySet()) { + versions.put( + entry.getKey(), + Collections.singletonList(entry.getValue())); + } + return new CanonicalSnapshot( + new TemporalState(versions, Collections.emptyList()), + accepted, + generatingEventIds); + } + + private static void apply( + NormalizedMemoryEvent event, + Map current, + Map generatingEventIds) { + FactKey key = event.getFactKey(); + MemoryEventOperation operation = event.getOperation(); + if (operation == MemoryEventOperation.RETRACT) { + MemoryFactVersion removed = current.remove(key); + if (removed == null) { + throw new IllegalArgumentException( + "Cannot retract a fact without current state: " + + event.getFactId()); + } + generatingEventIds.remove(removed.getId()); + return; + } + if (operation == MemoryEventOperation.CORRECT + && !current.containsKey(key)) { + throw new IllegalArgumentException( + "Cannot correct a fact without current state: " + + event.getFactId()); + } + if (operation != MemoryEventOperation.ADD + && operation != MemoryEventOperation.CORRECT) { + throw new UnsupportedOperationException( + "Unsupported memory event operation: " + operation); + } + + MemoryFactVersion version = new MemoryFactVersion( + event.getEventId() + ":version:0", + event.getEvent().getFact().get(), + MemoryFactVersionStatus.ACTIVE, + ALL_VALID_TIME, + TimeInterval.unboundedFrom(event.getRecordedAt()), + event.getEvidence()); + MemoryFactVersion replaced = current.put(key, version); + if (replaced != null) { + generatingEventIds.remove(replaced.getId()); + } + generatingEventIds.put(version.getId(), event.getEventId()); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/baseline/SingleTimestampBaseline.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/baseline/SingleTimestampBaseline.java new file mode 100644 index 000000000..c58d23efc --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/baseline/SingleTimestampBaseline.java @@ -0,0 +1,169 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.baseline; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.TreeMap; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.oracle.ReplayMethod; +import org.apache.geaflow.ai.temporal.semantics.CanonicalSnapshot; +import org.apache.geaflow.ai.temporal.semantics.EventLedger; +import org.apache.geaflow.ai.temporal.semantics.EventLedgerDecision; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; + +/** + * Keeps the current distinct values while ignoring valid-time intervals. + */ +public final class SingleTimestampBaseline implements ReplayMethod { + + private static final Comparator EVENT_ORDER = + Comparator.comparing(NormalizedMemoryEvent::getRecordedAt) + .thenComparing(NormalizedMemoryEvent::getEventId); + + private static final TimeInterval ALL_VALID_TIME = + TimeInterval.unboundedFrom(Instant.ofEpochMilli(Long.MIN_VALUE)); + + @Override + public CanonicalSnapshot replayToSnapshot( + List events) { + List ordered = new ArrayList<>( + Objects.requireNonNull(events, "events")); + for (NormalizedMemoryEvent event : ordered) { + Objects.requireNonNull(event, "event"); + } + Collections.sort(ordered, EVENT_ORDER); + + EventLedger ledger = new EventLedger(); + List accepted = new ArrayList<>(); + Map> current = + new TreeMap<>(); + for (NormalizedMemoryEvent event : ordered) { + EventLedgerDecision decision = ledger.check(event); + if (decision == EventLedgerDecision.DUPLICATE_NOOP) { + continue; + } + if (decision == EventLedgerDecision.REJECT_EVENT_ID_REUSE) { + throw new IllegalArgumentException( + "Event id reused with a different payload: " + + event.getEventId()); + } + + Map next = new TreeMap<>(); + Map existing = + current.get(event.getFactKey()); + if (existing != null) { + next.putAll(existing); + } + apply(event, next); + ledger.commit(event); + if (next.isEmpty()) { + current.remove(event.getFactKey()); + } else { + current.put(event.getFactKey(), next); + } + accepted.add(event); + } + + Map> versions = + new TreeMap<>(); + Map generatingEventIds = new TreeMap<>(); + List relations = new ArrayList<>(); + for (Map.Entry> entry + : current.entrySet()) { + List versionsForKey = new ArrayList<>(); + for (NormalizedMemoryEvent event : entry.getValue().values()) { + MemoryFactVersion version = version(event); + versionsForKey.add(version); + generatingEventIds.put( + version.getId(), event.getEventId()); + } + versions.put(entry.getKey(), versionsForKey); + addConflicts(versionsForKey, relations); + } + return new CanonicalSnapshot( + new TemporalState(versions, relations), + accepted, + generatingEventIds); + } + + private static void apply( + NormalizedMemoryEvent event, + Map current) { + MemoryEventOperation operation = event.getOperation(); + if (operation == MemoryEventOperation.ADD) { + current.put(event.getFactValue().get(), event); + return; + } + if (current.isEmpty()) { + throw new IllegalArgumentException( + "Operation requires current values for fact key: " + + event.getFactKey()); + } + current.clear(); + if (operation == MemoryEventOperation.CORRECT) { + current.put(event.getFactValue().get(), event); + } else if (operation != MemoryEventOperation.RETRACT) { + throw new UnsupportedOperationException( + "Unsupported memory event operation: " + operation); + } + } + + private static MemoryFactVersion version( + NormalizedMemoryEvent event) { + return new MemoryFactVersion( + versionId(event), + event.getEvent().getFact().get(), + MemoryFactVersionStatus.ACTIVE, + ALL_VALID_TIME, + TimeInterval.unboundedFrom(event.getRecordedAt()), + event.getEvidence()); + } + + private static String versionId(NormalizedMemoryEvent event) { + return event.getEventId() + ":version:0"; + } + + private static void addConflicts( + List versions, + List relations) { + for (int left = 0; left < versions.size(); left++) { + for (int right = left + 1; right < versions.size(); right++) { + relations.add(new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + versions.get(left).getId(), + versions.get(right).getId())); + } + } + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalIntegrator.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalIntegrator.java new file mode 100644 index 000000000..5a62af740 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalIntegrator.java @@ -0,0 +1,716 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.TreeMap; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.semantics.EventLedger; +import org.apache.geaflow.ai.temporal.semantics.EventLedgerDecision; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; + +/** + * Incrementally integrates legacy events by fact id and normalized events by + * fact key. + */ +public final class IncrementalTemporalIntegrator { + + private static final Comparator EVENT_ORDER = + Comparator.comparing(MemoryEvent::getTransactionTime) + .thenComparing(MemoryEvent::getId); + + private static final Comparator + NORMALIZED_EVENT_ORDER = + Comparator.comparing(NormalizedMemoryEvent::getRecordedAt) + .thenComparing(NormalizedMemoryEvent::getEventId); + + private static final Comparator VERSION_ORDER = + Comparator.comparing( + (MemoryFactVersion version) -> + version.getTransactionTime().getStart()) + .thenComparing( + version -> version.getValidTime().getStart()) + .thenComparing(MemoryFactVersion::getId); + + private static final Comparator VALID_TIME_ORDER = + Comparator.comparing( + (MemoryFactVersion version) -> + version.getValidTime().getStart()) + .thenComparing(MemoryFactVersion::getId); + + private static final Comparator EVIDENCE_ORDER = + Comparator.comparing(Evidence::getId) + .thenComparing(evidence -> evidence.getSource().getId()) + .thenComparing(evidence -> evidence.getSource().getName()) + .thenComparing(Evidence::getContent); + + private final Map eventsById = new HashMap<>(); + private final Map> eventsByFactId = + new HashMap<>(); + private final Map> versionsByFactId = + new HashMap<>(); + + private final EventLedger normalizedEventLedger = new EventLedger(); + private final Map> + normalizedEventsByFactKey = new HashMap<>(); + private final Map> + normalizedVersionsByFactKey = new HashMap<>(); + private final Map> + normalizedRelationsByFactKey = new HashMap<>(); + + public void apply(MemoryEvent event) { + Objects.requireNonNull(event, "event"); + + MemoryEvent existing = eventsById.get(event.getId()); + if (existing != null) { + if (!existing.equals(event)) { + throw new IllegalArgumentException( + "Conflicting event id: " + event.getId()); + } + return; + } + + String factId = event.getFactId(); + List existingEvents = eventsByFactId.get(factId); + List updatedEvents = existingEvents == null + ? new ArrayList<>() : new ArrayList<>(existingEvents); + boolean appended = existingEvents == null + || EVENT_ORDER.compare( + existingEvents.get(existingEvents.size() - 1), + event) < 0; + + updatedEvents.add(event); + Collections.sort(updatedEvents, EVENT_ORDER); + + List updatedVersions; + if (appended) { + List existingVersions = + versionsByFactId.get(factId); + updatedVersions = existingVersions == null + ? new ArrayList<>() + : new ArrayList<>(existingVersions); + applyOrderedEvent(event, updatedVersions); + } else { + updatedVersions = replayFact(updatedEvents); + } + + Collections.sort(updatedVersions, VERSION_ORDER); + eventsById.put(event.getId(), event); + eventsByFactId.put( + factId, + Collections.unmodifiableList(updatedEvents)); + versionsByFactId.put( + factId, + Collections.unmodifiableList(updatedVersions)); + } + + /** + * Applies an already normalized event to the state for its fact key. + */ + public void apply(NormalizedMemoryEvent event) { + Objects.requireNonNull(event, "event"); + + EventLedgerDecision decision = normalizedEventLedger.check(event); + if (decision == EventLedgerDecision.DUPLICATE_NOOP) { + return; + } + if (decision == EventLedgerDecision.REJECT_EVENT_ID_REUSE) { + throw new IllegalArgumentException( + "Event id reused with a different payload: " + + event.getEventId()); + } + + FactKey factKey = event.getFactKey(); + List existingEvents = + normalizedEventsByFactKey.get(factKey); + List updatedEvents = + existingEvents == null + ? new ArrayList<>() + : new ArrayList<>(existingEvents); + boolean appended = existingEvents == null + || NORMALIZED_EVENT_ORDER.compare( + existingEvents.get(existingEvents.size() - 1), + event) < 0; + updatedEvents.add(event); + Collections.sort(updatedEvents, NORMALIZED_EVENT_ORDER); + + List updatedVersions; + List updatedRelations; + if (appended) { + List existingVersions = + normalizedVersionsByFactKey.get(factKey); + List existingRelations = + normalizedRelationsByFactKey.get(factKey); + updatedVersions = existingVersions == null + ? new ArrayList<>() + : new ArrayList<>(existingVersions); + updatedRelations = existingRelations == null + ? new ArrayList<>() + : new ArrayList<>(existingRelations); + applyNormalizedOrderedEvent( + event, + updatedVersions, + updatedRelations); + } else { + updatedVersions = new ArrayList<>(); + updatedRelations = new ArrayList<>(); + for (NormalizedMemoryEvent orderedEvent : updatedEvents) { + applyNormalizedOrderedEvent( + orderedEvent, + updatedVersions, + updatedRelations); + } + } + + Collections.sort(updatedVersions, VERSION_ORDER); + Collections.sort(updatedRelations); + normalizedEventLedger.commit(event); + normalizedEventsByFactKey.put( + factKey, + Collections.unmodifiableList(updatedEvents)); + normalizedVersionsByFactKey.put( + factKey, + Collections.unmodifiableList(updatedVersions)); + normalizedRelationsByFactKey.put( + factKey, + Collections.unmodifiableList(updatedRelations)); + } + + public List snapshot() { + List snapshot = new ArrayList<>(); + for (List versions : + versionsByFactId.values()) { + snapshot.addAll(versions); + } + + Collections.sort(snapshot, VERSION_ORDER); + return Collections.unmodifiableList(snapshot); + } + + /** + * Returns an immutable, deterministic snapshot of normalized state. + */ + public TemporalState stateSnapshot() { + Map> versionsByFactKey = + new TreeMap<>(); + versionsByFactKey.putAll(normalizedVersionsByFactKey); + + List relations = new ArrayList<>(); + for (List keyRelations : + normalizedRelationsByFactKey.values()) { + relations.addAll(keyRelations); + } + Collections.sort(relations); + return new TemporalState(versionsByFactKey, relations); + } + + List eventSnapshot() { + List snapshot = + new ArrayList<>(eventsById.values()); + Collections.sort(snapshot, EVENT_ORDER); + return Collections.unmodifiableList(snapshot); + } + + private static void applyNormalizedOrderedEvent( + NormalizedMemoryEvent event, + List versions, + List relations) { + MemoryEventOperation operation = event.getOperation(); + if (operation == MemoryEventOperation.ADD) { + applyNormalizedAdd(event, versions, relations); + } else if (operation == MemoryEventOperation.CORRECT + || operation == MemoryEventOperation.RETRACT) { + applyNormalizedChange(event, versions, relations); + } else { + throw new UnsupportedOperationException( + "Unsupported memory event operation: " + + operation); + } + } + + private static void applyNormalizedAdd( + NormalizedMemoryEvent event, + List versions, + List relations) { + List duplicates = new ArrayList<>(); + for (MemoryFactVersion version : versions) { + if (isCurrentActive(version) + && version.getValidTime().overlaps( + event.getValidTime()) + && hasSameValue(event, version)) { + duplicates.add(version); + } + } + Collections.sort(duplicates, VALID_TIME_ORDER); + + TimeInterval validTime = event.getValidTime(); + List evidence = new ArrayList<>(); + mergeEvidence(evidence, event.getEvidence()); + Map materializedDuplicates = new HashMap<>(); + for (MemoryFactVersion duplicate : duplicates) { + versions.remove(duplicate); + boolean materialized = materializeClosedVersion( + duplicate, + event.getRecordedAt(), + versions, + relations); + materializedDuplicates.put(duplicate.getId(), materialized); + validTime = span(validTime, duplicate.getValidTime()); + mergeEvidence(evidence, duplicate.getEvidence()); + } + + MemoryFactVersion added = new MemoryFactVersion( + event.getEventId() + ":version:0", + event.getEvent().getFact().get(), + MemoryFactVersionStatus.ACTIVE, + validTime, + TimeInterval.unboundedFrom(event.getRecordedAt()), + evidence); + versions.add(added); + + for (MemoryFactVersion duplicate : duplicates) { + if (Boolean.TRUE.equals( + materializedDuplicates.get(duplicate.getId()))) { + addRelation( + relations, + new VersionRelation( + VersionRelationType.DUPLICATE_OF, + added.getId(), + duplicate.getId())); + } + } + + for (MemoryFactVersion version : versions) { + if (version != added + && isCurrentActive(version) + && version.getValidTime().overlaps(validTime) + && !hasSameValue(event, version)) { + addRelation( + relations, + new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + added.getId(), + version.getId())); + } + } + } + + private static void applyNormalizedChange( + NormalizedMemoryEvent event, + List versions, + List relations) { + List affected = new ArrayList<>(); + for (MemoryFactVersion version : versions) { + if (isCurrentActive(version) + && version.getValidTime().overlaps( + event.getValidTime())) { + affected.add(version); + } + } + Collections.sort(affected, VALID_TIME_ORDER); + if (!isFullyCovered(event.getValidTime(), affected)) { + throw new IllegalArgumentException( + "Event interval is not fully covered for fact key: " + + event.getFactKey()); + } + + Map materializedSources = new HashMap<>(); + for (MemoryFactVersion version : affected) { + versions.remove(version); + materializedSources.put( + version.getId(), + materializeClosedVersion( + version, + event.getRecordedAt(), + versions, + relations)); + } + + int versionIndex; + if (event.getOperation() == MemoryEventOperation.CORRECT) { + MemoryFactVersion corrected = new MemoryFactVersion( + event.getEventId() + ":version:0", + event.getEvent().getFact().get(), + MemoryFactVersionStatus.ACTIVE, + event.getValidTime(), + TimeInterval.unboundedFrom(event.getRecordedAt()), + event.getEvidence()); + versions.add(corrected); + for (MemoryFactVersion source : affected) { + addSupersedesIfMaterialized( + corrected, + source, + materializedSources, + relations); + } + versionIndex = 1; + } else { + versionIndex = addTombstones( + event, + affected, + materializedSources, + versions, + relations); + } + + for (MemoryFactVersion source : affected) { + for (TimeInterval remaining : source.getValidTime() + .subtract(event.getValidTime())) { + MemoryFactVersion residue = new MemoryFactVersion( + event.getEventId() + ":version:" + + versionIndex++, + source.getFact(), + MemoryFactVersionStatus.ACTIVE, + remaining, + TimeInterval.unboundedFrom( + event.getRecordedAt()), + source.getEvidence()); + versions.add(residue); + addSupersedesIfMaterialized( + residue, + source, + materializedSources, + relations); + } + } + + addCurrentConflicts(versions, relations); + } + + private static void addCurrentConflicts( + List versions, + List relations) { + for (int leftIndex = 0; + leftIndex < versions.size(); leftIndex++) { + MemoryFactVersion left = versions.get(leftIndex); + if (!isCurrentActive(left)) { + continue; + } + for (int rightIndex = leftIndex + 1; + rightIndex < versions.size(); rightIndex++) { + MemoryFactVersion right = versions.get(rightIndex); + if (isCurrentActive(right) + && left.getValidTime().overlaps( + right.getValidTime()) + && !factValue(left.getFact()).equals( + factValue(right.getFact()))) { + addRelation( + relations, + new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + left.getId(), + right.getId())); + } + } + } + } + + private static int addTombstones( + NormalizedMemoryEvent event, + List affected, + Map materializedSources, + List versions, + List relations) { + int versionIndex = 0; + for (MemoryFactVersion source : affected) { + TimeInterval overlap = source.getValidTime() + .intersection(event.getValidTime()).get(); + MemoryFactVersion tombstone = new MemoryFactVersion( + event.getEventId() + ":version:" + versionIndex++, + source.getFact(), + MemoryFactVersionStatus.RETRACTED, + overlap, + TimeInterval.unboundedFrom(event.getRecordedAt()), + event.getEvidence()); + versions.add(tombstone); + addSupersedesIfMaterialized( + tombstone, + source, + materializedSources, + relations); + } + return versionIndex; + } + + private static boolean materializeClosedVersion( + MemoryFactVersion version, + Instant recordedAt, + List versions, + List relations) { + if (!version.getTransactionTime().getStart() + .isBefore(recordedAt)) { + removeRelationsFor(version.getId(), relations); + return false; + } + + versions.add(new MemoryFactVersion( + version.getId(), + version.getFact(), + version.getStatus(), + version.getValidTime(), + new TimeInterval( + version.getTransactionTime().getStart(), + recordedAt), + version.getEvidence())); + return true; + } + + private static void addSupersedesIfMaterialized( + MemoryFactVersion replacement, + MemoryFactVersion source, + Map materializedSources, + List relations) { + if (Boolean.TRUE.equals( + materializedSources.get(source.getId()))) { + addRelation( + relations, + new VersionRelation( + VersionRelationType.SUPERSEDES, + replacement.getId(), + source.getId())); + } + } + + private static void addRelation( + List relations, + VersionRelation relation) { + if (!relations.contains(relation)) { + relations.add(relation); + } + } + + private static void removeRelationsFor( + String versionId, + List relations) { + for (int index = relations.size() - 1; index >= 0; index--) { + VersionRelation relation = relations.get(index); + if (relation.getFromVersionId().equals(versionId) + || relation.getToVersionId().equals(versionId)) { + relations.remove(index); + } + } + } + + private static boolean hasSameValue( + NormalizedMemoryEvent event, + MemoryFactVersion version) { + Optional eventValue = event.getFactValue(); + return eventValue.isPresent() + && eventValue.get().equals(factValue(version.getFact())); + } + + private static FactValue factValue(MemoryFact fact) { + if (fact.isRelationship()) { + return FactValue.entityReference( + fact.getTarget().get().getId()); + } + return FactValue.literal(fact.getLiteralValue().get()); + } + + private static TimeInterval span( + TimeInterval left, + TimeInterval right) { + Instant start = left.getStart().isBefore(right.getStart()) + ? left.getStart() : right.getStart(); + Instant end; + if (!left.getEnd().isPresent() + || !right.getEnd().isPresent()) { + end = null; + } else { + Instant leftEnd = left.getEnd().get(); + Instant rightEnd = right.getEnd().get(); + end = leftEnd.isAfter(rightEnd) ? leftEnd : rightEnd; + } + return new TimeInterval(start, end); + } + + private static void mergeEvidence( + List target, + List additions) { + for (Evidence evidence : additions) { + if (!target.contains(evidence)) { + target.add(evidence); + } + } + Collections.sort(target, EVIDENCE_ORDER); + } + + private static boolean isCurrentActive( + MemoryFactVersion version) { + return version.getStatus() == MemoryFactVersionStatus.ACTIVE + && isCurrent(version); + } + + private static List replayFact( + List events) { + List versions = new ArrayList<>(); + for (MemoryEvent event : events) { + applyOrderedEvent(event, versions); + } + return versions; + } + + private static void applyOrderedEvent( + MemoryEvent event, + List versions) { + MemoryEventOperation operation = event.getOperation(); + if (operation == MemoryEventOperation.ADD) { + applyAdd(event, versions); + } else if (operation == MemoryEventOperation.CORRECT + || operation == MemoryEventOperation.RETRACT) { + applyChange(event, versions); + } else { + throw new UnsupportedOperationException( + "Unsupported memory event operation: " + + operation); + } + } + + private static void applyAdd( + MemoryEvent event, + List versions) { + for (MemoryFactVersion version : versions) { + if (isCurrent(version) + && version.getFact().getId().equals(event.getFactId()) + && version.getValidTime().overlaps( + event.getValidTime())) { + throw new IllegalArgumentException( + "Overlapping add for fact id: " + + event.getFactId()); + } + } + + versions.add(new MemoryFactVersion( + event.getId() + ":version:0", + event.getFact().get(), + event.getValidTime(), + TimeInterval.unboundedFrom( + event.getTransactionTime()), + event.getEvidence())); + } + + private static void applyChange( + MemoryEvent event, + List versions) { + List affected = new ArrayList<>(); + for (MemoryFactVersion version : versions) { + if (isCurrent(version) + && version.getFact().getId().equals(event.getFactId()) + && version.getValidTime().overlaps( + event.getValidTime())) { + affected.add(version); + } + } + + Collections.sort(affected, VALID_TIME_ORDER); + if (!isFullyCovered(event.getValidTime(), affected)) { + throw new IllegalArgumentException( + "Event interval is not fully covered for fact id: " + + event.getFactId()); + } + + versions.removeAll(affected); + + int fragmentIndex = 1; + for (MemoryFactVersion version : affected) { + if (version.getTransactionTime().getStart() + .isBefore(event.getTransactionTime())) { + versions.add(new MemoryFactVersion( + version.getId(), + version.getFact(), + version.getValidTime(), + new TimeInterval( + version.getTransactionTime().getStart(), + event.getTransactionTime()), + version.getEvidence())); + } + + for (TimeInterval remaining : + version.getValidTime().subtract( + event.getValidTime())) { + versions.add(new MemoryFactVersion( + event.getId() + ":version:" + + fragmentIndex++, + version.getFact(), + remaining, + TimeInterval.unboundedFrom( + event.getTransactionTime()), + version.getEvidence())); + } + } + + if (event.getOperation() == MemoryEventOperation.CORRECT) { + versions.add(new MemoryFactVersion( + event.getId() + ":version:0", + event.getFact().get(), + event.getValidTime(), + TimeInterval.unboundedFrom( + event.getTransactionTime()), + event.getEvidence())); + } + } + + private static boolean isFullyCovered( + TimeInterval target, + List coveringVersions) { + List uncovered = new ArrayList<>(); + uncovered.add(target); + + for (MemoryFactVersion version : coveringVersions) { + List remaining = new ArrayList<>(); + for (TimeInterval interval : uncovered) { + remaining.addAll( + interval.subtract(version.getValidTime())); + } + + uncovered = remaining; + if (uncovered.isEmpty()) { + return true; + } + } + + return false; + } + + private static boolean isCurrent( + MemoryFactVersion version) { + return !version.getTransactionTime() + .getEnd().isPresent(); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/integration/TemporalEventAggregateFunction.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/integration/TemporalEventAggregateFunction.java new file mode 100644 index 000000000..dc37bda5d --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/integration/TemporalEventAggregateFunction.java @@ -0,0 +1,73 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.util.List; +import java.util.Objects; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.api.function.base.AggregateFunction; + +/** + * Adapts temporal event integration to GeaFlow keyed aggregation. + */ +public final class TemporalEventAggregateFunction implements + AggregateFunction> { + + @Override + public IncrementalTemporalIntegrator createAccumulator() { + return new IncrementalTemporalIntegrator(); + } + + @Override + public void add( + MemoryEvent value, + IncrementalTemporalIntegrator accumulator) { + Objects.requireNonNull(accumulator, "accumulator") + .apply(value); + } + + @Override + public List getResult( + IncrementalTemporalIntegrator accumulator) { + return Objects.requireNonNull( + accumulator, + "accumulator").snapshot(); + } + + @Override + public IncrementalTemporalIntegrator merge( + IncrementalTemporalIntegrator left, + IncrementalTemporalIntegrator right) { + Objects.requireNonNull(left, "left"); + Objects.requireNonNull(right, "right"); + + IncrementalTemporalIntegrator merged = + new IncrementalTemporalIntegrator(); + for (MemoryEvent event : left.eventSnapshot()) { + merged.apply(event); + } + for (MemoryEvent event : right.eventSnapshot()) { + merged.apply(event); + } + return merged; + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/Evidence.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/Evidence.java new file mode 100644 index 000000000..998137fda --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/Evidence.java @@ -0,0 +1,78 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.Objects; + +/** + * An immutable piece of evidence and its source. + */ +public final class Evidence { + + private final String id; + private final Source source; + private final String content; + + public Evidence(String id, Source source, String content) { + this.id = requireText(id, "id"); + this.source = Objects.requireNonNull(source, "source"); + this.content = requireText(content, "content"); + } + + public String getId() { + return id; + } + + public Source getSource() { + return source; + } + + public String getContent() { + return content; + } + + private static String requireText(String value, String fieldName) { + Objects.requireNonNull(value, fieldName); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Evidence " + fieldName + " must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof Evidence)) { + return false; + } + Evidence that = (Evidence) object; + return id.equals(that.id) + && source.equals(that.source) + && content.equals(that.content); + } + + @Override + public int hashCode() { + return Objects.hash(id, source, content); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/FactKey.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/FactKey.java new file mode 100644 index 000000000..b74621b37 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/FactKey.java @@ -0,0 +1,96 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.Objects; + +/** + * An immutable identity for a fact independent of its value. + */ +public final class FactKey implements Comparable { + + private final String subjectId; + private final String predicate; + private final String scope; + + public FactKey( + String subjectId, + String predicate, + String scope) { + this.subjectId = requireText(subjectId, "subjectId"); + this.predicate = requireText(predicate, "predicate"); + this.scope = requireText(scope, "scope"); + } + + public String getSubjectId() { + return subjectId; + } + + public String getPredicate() { + return predicate; + } + + public String getScope() { + return scope; + } + + @Override + public int compareTo(FactKey that) { + int comparison = subjectId.compareTo(that.subjectId); + if (comparison != 0) { + return comparison; + } + comparison = predicate.compareTo(that.predicate); + if (comparison != 0) { + return comparison; + } + return scope.compareTo(that.scope); + } + + private static String requireText( + String value, + String fieldName) { + Objects.requireNonNull(value, fieldName); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Fact key " + fieldName + " must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof FactKey)) { + return false; + } + FactKey that = (FactKey) object; + return subjectId.equals(that.subjectId) + && predicate.equals(that.predicate) + && scope.equals(that.scope); + } + + @Override + public int hashCode() { + return Objects.hash(subjectId, predicate, scope); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/FactValue.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/FactValue.java new file mode 100644 index 000000000..5905f82ea --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/FactValue.java @@ -0,0 +1,101 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.Objects; +import java.util.Optional; + +/** + * An immutable literal or entity-reference fact value. + */ +public final class FactValue implements Comparable { + + public enum Kind { + LITERAL, + ENTITY_REF + } + + private final Kind kind; + private final String value; + + private FactValue(Kind kind, String value) { + this.kind = Objects.requireNonNull(kind, "kind"); + this.value = requireText(value); + } + + public static FactValue literal(String value) { + return new FactValue(Kind.LITERAL, value); + } + + public static FactValue entityReference(String entityId) { + return new FactValue(Kind.ENTITY_REF, entityId); + } + + public Kind getKind() { + return kind; + } + + public String getValue() { + return value; + } + + public Optional getLiteralValue() { + return kind == Kind.LITERAL + ? Optional.of(value) : Optional.empty(); + } + + public Optional getEntityId() { + return kind == Kind.ENTITY_REF + ? Optional.of(value) : Optional.empty(); + } + + @Override + public int compareTo(FactValue that) { + int comparison = kind.compareTo(that.kind); + return comparison != 0 + ? comparison : value.compareTo(that.value); + } + + private static String requireText(String value) { + Objects.requireNonNull(value, "value"); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Fact value must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof FactValue)) { + return false; + } + FactValue that = (FactValue) object; + return kind == that.kind && value.equals(that.value); + } + + @Override + public int hashCode() { + return Objects.hash(kind, value); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEntity.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEntity.java new file mode 100644 index 000000000..87ce52a8c --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEntity.java @@ -0,0 +1,70 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.Objects; + +/** + * An immutable identity anchor for temporal memory facts. + */ +public final class MemoryEntity { + + private final String id; + private final String label; + + public MemoryEntity(String id, String label) { + this.id = requireText(id, "id"); + this.label = requireText(label, "label"); + } + + public String getId() { + return id; + } + + public String getLabel() { + return label; + } + + private static String requireText(String value, String fieldName) { + Objects.requireNonNull(value, fieldName); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Memory entity " + fieldName + " must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof MemoryEntity)) { + return false; + } + MemoryEntity that = (MemoryEntity) object; + return id.equals(that.id) && label.equals(that.label); + } + + @Override + public int hashCode() { + return Objects.hash(id, label); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEvent.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEvent.java new file mode 100644 index 000000000..122cee8a3 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEvent.java @@ -0,0 +1,194 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Objects; +import java.util.Optional; + +/** + * An immutable input event for temporal memory replay. + */ +public final class MemoryEvent { + + private final String id; + private final MemoryEventOperation operation; + private final String factId; + private final MemoryFact fact; + private final TimeInterval validTime; + private final Instant transactionTime; + private final List evidence; + + private MemoryEvent( + String id, + MemoryEventOperation operation, + String factId, + MemoryFact fact, + TimeInterval validTime, + Instant transactionTime, + List evidence) { + this.id = requireText(id, "id"); + this.operation = + Objects.requireNonNull(operation, "operation"); + this.factId = requireText(factId, "fact id"); + this.fact = fact; + this.validTime = + Objects.requireNonNull(validTime, "validTime"); + this.transactionTime = + Objects.requireNonNull(transactionTime, "transactionTime"); + this.evidence = copyEvidence(evidence); + } + + public static MemoryEvent add( + String id, + MemoryFact fact, + TimeInterval validTime, + Instant transactionTime, + List evidence) { + Objects.requireNonNull(fact, "fact"); + return new MemoryEvent( + id, + MemoryEventOperation.ADD, + fact.getId(), + fact, + validTime, + transactionTime, + evidence); + } + + public static MemoryEvent correct( + String id, + MemoryFact fact, + TimeInterval validTime, + Instant transactionTime, + List evidence) { + Objects.requireNonNull(fact, "fact"); + return new MemoryEvent( + id, + MemoryEventOperation.CORRECT, + fact.getId(), + fact, + validTime, + transactionTime, + evidence); + } + + public static MemoryEvent retract( + String id, + String factId, + TimeInterval validTime, + Instant transactionTime, + List evidence) { + return new MemoryEvent( + id, + MemoryEventOperation.RETRACT, + factId, + null, + validTime, + transactionTime, + evidence); + } + + public String getId() { + return id; + } + + public MemoryEventOperation getOperation() { + return operation; + } + + public String getFactId() { + return factId; + } + + public Optional getFact() { + return Optional.ofNullable(fact); + } + + public TimeInterval getValidTime() { + return validTime; + } + + public Instant getTransactionTime() { + return transactionTime; + } + + public List getEvidence() { + return evidence; + } + + private static List copyEvidence( + List evidence) { + List copy = + new ArrayList<>(Objects.requireNonNull(evidence, "evidence")); + if (copy.isEmpty()) { + throw new IllegalArgumentException( + "Memory event evidence must not be empty"); + } + for (Evidence item : copy) { + Objects.requireNonNull(item, "evidence item"); + } + return Collections.unmodifiableList(copy); + } + + private static String requireText( + String value, + String fieldName) { + Objects.requireNonNull(value, fieldName); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Memory event " + fieldName + " must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof MemoryEvent)) { + return false; + } + MemoryEvent that = (MemoryEvent) object; + return id.equals(that.id) + && operation == that.operation + && factId.equals(that.factId) + && Objects.equals(fact, that.fact) + && validTime.equals(that.validTime) + && transactionTime.equals(that.transactionTime) + && evidence.equals(that.evidence); + } + + @Override + public int hashCode() { + return Objects.hash( + id, + operation, + factId, + fact, + validTime, + transactionTime, + evidence); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEventOperation.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEventOperation.java new file mode 100644 index 000000000..7e5d63460 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryEventOperation.java @@ -0,0 +1,30 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +/** + * Supported operations for temporal memory events. + */ +public enum MemoryEventOperation { + + ADD, + CORRECT, + RETRACT +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFact.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFact.java new file mode 100644 index 000000000..23ead3e59 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFact.java @@ -0,0 +1,133 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.Objects; +import java.util.Optional; + +/** + * An immutable structured assertion about a memory entity. + */ +public final class MemoryFact { + + private final String id; + private final MemoryEntity subject; + private final String predicate; + private final String literalValue; + private final MemoryEntity target; + + private MemoryFact( + String id, + MemoryEntity subject, + String predicate, + String literalValue, + MemoryEntity target) { + this.id = requireText(id, "id"); + this.subject = Objects.requireNonNull(subject, "subject"); + this.predicate = requireText(predicate, "predicate"); + this.literalValue = literalValue; + this.target = target; + } + + public static MemoryFact attribute( + String id, + MemoryEntity subject, + String predicate, + String literalValue) { + return new MemoryFact( + id, + subject, + predicate, + requireText(literalValue, "literal value"), + null); + } + + public static MemoryFact relationship( + String id, + MemoryEntity subject, + String predicate, + MemoryEntity target) { + return new MemoryFact( + id, + subject, + predicate, + null, + Objects.requireNonNull(target, "target")); + } + + public String getId() { + return id; + } + + public MemoryEntity getSubject() { + return subject; + } + + public String getPredicate() { + return predicate; + } + + public Optional getLiteralValue() { + return Optional.ofNullable(literalValue); + } + + public Optional getTarget() { + return Optional.ofNullable(target); + } + + public boolean isRelationship() { + return target != null; + } + + private static String requireText(String value, String fieldName) { + Objects.requireNonNull(value, fieldName); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Memory fact " + fieldName + " must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof MemoryFact)) { + return false; + } + MemoryFact that = (MemoryFact) object; + return id.equals(that.id) + && subject.equals(that.subject) + && predicate.equals(that.predicate) + && Objects.equals(literalValue, that.literalValue) + && Objects.equals(target, that.target); + } + + @Override + public int hashCode() { + return Objects.hash( + id, + subject, + predicate, + literalValue, + target); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersion.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersion.java new file mode 100644 index 000000000..a52de766f --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersion.java @@ -0,0 +1,140 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Objects; + +/** + * An immutable bitemporal version of a memory fact. + */ +public final class MemoryFactVersion { + + private final String id; + private final MemoryFact fact; + private final MemoryFactVersionStatus status; + private final TimeInterval validTime; + private final TimeInterval transactionTime; + private final List evidence; + + public MemoryFactVersion( + String id, + MemoryFact fact, + TimeInterval validTime, + TimeInterval transactionTime, + List evidence) { + this( + id, + fact, + MemoryFactVersionStatus.ACTIVE, + validTime, + transactionTime, + evidence); + } + + public MemoryFactVersion( + String id, + MemoryFact fact, + MemoryFactVersionStatus status, + TimeInterval validTime, + TimeInterval transactionTime, + List evidence) { + this.id = requireText(id); + this.fact = Objects.requireNonNull(fact, "fact"); + this.status = Objects.requireNonNull(status, "status"); + this.validTime = Objects.requireNonNull(validTime, "validTime"); + this.transactionTime = + Objects.requireNonNull(transactionTime, "transactionTime"); + + List evidenceCopy = + new ArrayList<>(Objects.requireNonNull(evidence, "evidence")); + if (evidenceCopy.isEmpty()) { + throw new IllegalArgumentException( + "Memory fact version evidence must not be empty"); + } + for (Evidence item : evidenceCopy) { + Objects.requireNonNull(item, "evidence item"); + } + this.evidence = Collections.unmodifiableList(evidenceCopy); + } + + public String getId() { + return id; + } + + public MemoryFact getFact() { + return fact; + } + + public MemoryFactVersionStatus getStatus() { + return status; + } + + public TimeInterval getValidTime() { + return validTime; + } + + public TimeInterval getTransactionTime() { + return transactionTime; + } + + public List getEvidence() { + return evidence; + } + + private static String requireText(String value) { + Objects.requireNonNull(value, "id"); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Memory fact version id must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof MemoryFactVersion)) { + return false; + } + MemoryFactVersion that = (MemoryFactVersion) object; + return id.equals(that.id) + && fact.equals(that.fact) + && status == that.status + && validTime.equals(that.validTime) + && transactionTime.equals(that.transactionTime) + && evidence.equals(that.evidence); + } + + @Override + public int hashCode() { + return Objects.hash( + id, + fact, + status, + validTime, + transactionTime, + evidence); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersionStatus.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersionStatus.java new file mode 100644 index 000000000..8e2e2e85d --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersionStatus.java @@ -0,0 +1,25 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +public enum MemoryFactVersionStatus { + ACTIVE, + RETRACTED +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/Source.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/Source.java new file mode 100644 index 000000000..c73576f86 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/Source.java @@ -0,0 +1,70 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.Objects; + +/** + * An immutable description of an evidence source. + */ +public final class Source { + + private final String id; + private final String name; + + public Source(String id, String name) { + this.id = requireText(id, "id"); + this.name = requireText(name, "name"); + } + + public String getId() { + return id; + } + + public String getName() { + return name; + } + + private static String requireText(String value, String fieldName) { + Objects.requireNonNull(value, fieldName); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Source " + fieldName + " must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof Source)) { + return false; + } + Source that = (Source) object; + return id.equals(that.id) && name.equals(that.name); + } + + @Override + public int hashCode() { + return Objects.hash(id, name); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/TimeInterval.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/TimeInterval.java new file mode 100644 index 000000000..bc372ca14 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/TimeInterval.java @@ -0,0 +1,137 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Objects; +import java.util.Optional; + +/** + * An immutable half-open time interval. + * + *

The start is inclusive and the end is exclusive. A null internal end + * represents positive infinity. + */ +public final class TimeInterval { + + private final Instant start; + private final Instant end; + + public TimeInterval(Instant start, Instant end) { + this.start = Objects.requireNonNull(start, "start"); + if (end != null && !start.isBefore(end)) { + throw new IllegalArgumentException( + "Interval start must be earlier than end"); + } + this.end = end; + } + + public static TimeInterval unboundedFrom(Instant start) { + return new TimeInterval(start, null); + } + + public Instant getStart() { + return start; + } + + public Optional getEnd() { + return Optional.ofNullable(end); + } + + public boolean contains(Instant time) { + Objects.requireNonNull(time, "time"); + return !time.isBefore(start) && (end == null || time.isBefore(end)); + } + + public boolean overlaps(TimeInterval other) { + return intersection(other).isPresent(); + } + + public Optional intersection(TimeInterval other) { + Objects.requireNonNull(other, "other"); + + Instant intersectionStart = + start.isAfter(other.start) ? start : other.start; + Instant intersectionEnd = earliestEnd(end, other.end); + + if (intersectionEnd != null + && !intersectionStart.isBefore(intersectionEnd)) { + return Optional.empty(); + } + return Optional.of( + new TimeInterval(intersectionStart, intersectionEnd)); + } + + public List subtract(TimeInterval other) { + Optional intersection = intersection(other); + if (!intersection.isPresent()) { + return Collections.singletonList(this); + } + + TimeInterval overlap = intersection.get(); + List remaining = new ArrayList<>(2); + + if (start.isBefore(overlap.start)) { + remaining.add(new TimeInterval(start, overlap.start)); + } + + if (overlap.end != null + && (end == null || overlap.end.isBefore(end))) { + remaining.add(new TimeInterval(overlap.end, end)); + } + + return remaining; + } + + private static Instant earliestEnd(Instant left, Instant right) { + if (left == null) { + return right; + } + if (right == null) { + return left; + } + return left.isBefore(right) ? left : right; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof TimeInterval)) { + return false; + } + TimeInterval that = (TimeInterval) object; + return start.equals(that.start) && Objects.equals(end, that.end); + } + + @Override + public int hashCode() { + return Objects.hash(start, end); + } + + @Override + public String toString() { + return "[" + start + ", " + (end == null ? "infinity" : end) + ")"; + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/VersionRelation.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/VersionRelation.java new file mode 100644 index 000000000..bf09076cb --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/VersionRelation.java @@ -0,0 +1,112 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.Objects; + +/** + * An immutable directed relation between two materialized versions. + */ +public final class VersionRelation + implements Comparable { + + private final VersionRelationType type; + private final String fromVersionId; + private final String toVersionId; + + public VersionRelation( + VersionRelationType type, + String fromVersionId, + String toVersionId) { + this.type = Objects.requireNonNull(type, "type"); + String from = requireText( + fromVersionId, + "fromVersionId"); + String to = requireText(toVersionId, "toVersionId"); + if (from.equals(to)) { + throw new IllegalArgumentException( + "Version relation must not be a self relation"); + } + if (type == VersionRelationType.CONFLICTS_WITH + && from.compareTo(to) > 0) { + this.fromVersionId = to; + this.toVersionId = from; + } else { + this.fromVersionId = from; + this.toVersionId = to; + } + } + + public VersionRelationType getType() { + return type; + } + + public String getFromVersionId() { + return fromVersionId; + } + + public String getToVersionId() { + return toVersionId; + } + + @Override + public int compareTo(VersionRelation that) { + int comparison = type.compareTo(that.type); + if (comparison != 0) { + return comparison; + } + comparison = fromVersionId.compareTo(that.fromVersionId); + if (comparison != 0) { + return comparison; + } + return toVersionId.compareTo(that.toVersionId); + } + + private static String requireText( + String value, + String fieldName) { + Objects.requireNonNull(value, fieldName); + if (value.trim().isEmpty()) { + throw new IllegalArgumentException( + "Version relation " + fieldName + + " must not be blank"); + } + return value; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof VersionRelation)) { + return false; + } + VersionRelation that = (VersionRelation) object; + return type == that.type + && fromVersionId.equals(that.fromVersionId) + && toVersionId.equals(that.toVersionId); + } + + @Override + public int hashCode() { + return Objects.hash(type, fromVersionId, toVersionId); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/VersionRelationType.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/VersionRelationType.java new file mode 100644 index 000000000..ab4ebbec6 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/model/VersionRelationType.java @@ -0,0 +1,26 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +public enum VersionRelationType { + SUPERSEDES, + DUPLICATE_OF, + CONFLICTS_WITH +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracle.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracle.java new file mode 100644 index 000000000..a78b7abb3 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracle.java @@ -0,0 +1,610 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.oracle; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Set; +import java.util.TreeSet; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.semantics.EventLedger; +import org.apache.geaflow.ai.temporal.semantics.EventLedgerDecision; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; + +/** + * Recomputes temporal memory versions from a complete event collection. + */ +public final class FullReplayOracle { + + private static final Comparator EVENT_ORDER = + Comparator.comparing(MemoryEvent::getTransactionTime) + .thenComparing(MemoryEvent::getId); + + private static final Comparator + NORMALIZED_EVENT_ORDER = + Comparator.comparing(NormalizedMemoryEvent::getRecordedAt) + .thenComparing(NormalizedMemoryEvent::getEventId); + + private static final Comparator VERSION_ORDER = + Comparator.comparing( + (MemoryFactVersion version) -> + version.getTransactionTime().getStart()) + .thenComparing( + version -> version.getValidTime().getStart()) + .thenComparing(MemoryFactVersion::getId); + + private static final Comparator EVIDENCE_ORDER = + Comparator.comparing(Evidence::getId) + .thenComparing(evidence -> evidence.getSource().getId()) + .thenComparing(evidence -> evidence.getSource().getName()) + .thenComparing(Evidence::getContent); + + private static final Comparator VALID_TIME_ORDER = + Comparator.comparing( + (MemoryFactVersion version) -> + version.getValidTime().getStart()) + .thenComparing(MemoryFactVersion::getId); + + public List replay( + List events) { + Objects.requireNonNull(events, "events"); + + Map uniqueEvents = new HashMap<>(); + for (MemoryEvent event : events) { + Objects.requireNonNull(event, "event"); + + MemoryEvent existing = uniqueEvents.get(event.getId()); + if (existing == null) { + uniqueEvents.put(event.getId(), event); + } else if (!existing.equals(event)) { + throw new IllegalArgumentException( + "Conflicting event id: " + event.getId()); + } + } + + List orderedEvents = + new ArrayList<>(uniqueEvents.values()); + Collections.sort(orderedEvents, EVENT_ORDER); + + List versions = new ArrayList<>(); + for (MemoryEvent event : orderedEvents) { + MemoryEventOperation operation = event.getOperation(); + if (operation == MemoryEventOperation.ADD) { + replayAdd(event, versions); + } else if (operation == MemoryEventOperation.CORRECT + || operation == MemoryEventOperation.RETRACT) { + replayChange(event, versions); + } else { + throw new UnsupportedOperationException( + "Unsupported memory event operation: " + + operation); + } + } + + Collections.sort(versions, VERSION_ORDER); + return Collections.unmodifiableList(versions); + } + + /** + * Recomputes canonical temporal state from normalized events. + */ + public TemporalState replayNormalized( + List events) { + Objects.requireNonNull(events, "events"); + + List orderedEvents = + new ArrayList<>(events.size()); + for (NormalizedMemoryEvent event : events) { + orderedEvents.add(Objects.requireNonNull(event, "event")); + } + Collections.sort(orderedEvents, NORMALIZED_EVENT_ORDER); + + EventLedger ledger = new EventLedger(); + Map> versionsByFactKey = + new HashMap<>(); + Set relations = new TreeSet<>(); + + for (NormalizedMemoryEvent event : orderedEvents) { + EventLedgerDecision decision = ledger.check(event); + if (decision == EventLedgerDecision.DUPLICATE_NOOP) { + continue; + } + if (decision + == EventLedgerDecision.REJECT_EVENT_ID_REUSE) { + throw new IllegalArgumentException( + "Event id reused with a different payload: " + + event.getEventId()); + } + + List nextVersions = + new ArrayList<>(versionsByFactKey.getOrDefault( + event.getFactKey(), + Collections.emptyList())); + Set nextRelations = + new TreeSet<>(relations); + replayNormalizedEvent( + event, + nextVersions, + nextRelations); + + ledger.commit(event); + versionsByFactKey.put( + event.getFactKey(), + nextVersions); + relations = nextRelations; + } + + for (List versions + : versionsByFactKey.values()) { + Collections.sort(versions, VERSION_ORDER); + } + return new TemporalState( + versionsByFactKey, + new ArrayList<>(relations)); + } + + private static void replayNormalizedEvent( + NormalizedMemoryEvent event, + List versions, + Set relations) { + if (event.getOperation() == MemoryEventOperation.ADD) { + replayNormalizedAdd(event, versions, relations); + return; + } + if (event.getOperation() == MemoryEventOperation.CORRECT + || event.getOperation() == MemoryEventOperation.RETRACT) { + replayNormalizedChange(event, versions, relations); + return; + } + throw new UnsupportedOperationException( + "Unsupported memory event operation: " + + event.getOperation()); + } + + private static void replayNormalizedAdd( + NormalizedMemoryEvent event, + List versions, + Set relations) { + FactValue newValue = event.getFactValue().get(); + TimeInterval mergedValidTime = event.getValidTime(); + List duplicates = new ArrayList<>(); + + boolean expanded; + do { + expanded = false; + for (MemoryFactVersion version : versions) { + if (isCurrentActive(version) + && !duplicates.contains(version) + && factValue(version.getFact()).equals(newValue) + && version.getValidTime().overlaps( + mergedValidTime)) { + duplicates.add(version); + mergedValidTime = span( + mergedValidTime, + version.getValidTime()); + expanded = true; + } + } + } while (expanded); + Collections.sort(duplicates, VALID_TIME_ORDER); + + List evidence = + new ArrayList<>(event.getEvidence()); + versions.removeAll(duplicates); + String newVersionId = versionId(event, 0); + for (MemoryFactVersion duplicate : duplicates) { + evidence.addAll(duplicate.getEvidence()); + if (closeVersion( + duplicate, + event.getRecordedAt(), + versions, + relations)) { + relations.add(new VersionRelation( + VersionRelationType.DUPLICATE_OF, + newVersionId, + duplicate.getId())); + } + } + + MemoryFactVersion added = new MemoryFactVersion( + newVersionId, + event.getEvent().getFact().get(), + MemoryFactVersionStatus.ACTIVE, + mergedValidTime, + TimeInterval.unboundedFrom(event.getRecordedAt()), + sortedUniqueEvidence(evidence)); + for (MemoryFactVersion version : versions) { + if (isCurrentActive(version) + && version.getValidTime().overlaps( + mergedValidTime) + && !factValue(version.getFact()).equals(newValue)) { + relations.add(new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + newVersionId, + version.getId())); + } + } + versions.add(added); + } + + private static void replayNormalizedChange( + NormalizedMemoryEvent event, + List versions, + Set relations) { + List affected = new ArrayList<>(); + for (MemoryFactVersion version : versions) { + if (isCurrentActive(version) + && version.getValidTime().overlaps( + event.getValidTime())) { + affected.add(version); + } + } + Collections.sort(affected, VALID_TIME_ORDER); + if (!isFullyCovered(event.getValidTime(), affected)) { + throw new IllegalArgumentException( + "Event interval is not fully covered for fact key: " + + event.getFactKey()); + } + + versions.removeAll(affected); + Set materializedSourceIds = new TreeSet<>(); + for (MemoryFactVersion version : affected) { + if (closeVersion( + version, + event.getRecordedAt(), + versions, + relations)) { + materializedSourceIds.add(version.getId()); + } + } + + int fragmentIndex; + if (event.getOperation() == MemoryEventOperation.CORRECT) { + MemoryFactVersion corrected = new MemoryFactVersion( + versionId(event, 0), + event.getEvent().getFact().get(), + MemoryFactVersionStatus.ACTIVE, + event.getValidTime(), + TimeInterval.unboundedFrom(event.getRecordedAt()), + event.getEvidence()); + versions.add(corrected); + addSupersedes( + corrected.getId(), + affected, + materializedSourceIds, + relations); + fragmentIndex = 1; + } else { + fragmentIndex = 0; + for (MemoryFactVersion source : affected) { + MemoryFactVersion tombstone = new MemoryFactVersion( + versionId(event, fragmentIndex++), + source.getFact(), + MemoryFactVersionStatus.RETRACTED, + source.getValidTime().intersection( + event.getValidTime()).get(), + TimeInterval.unboundedFrom( + event.getRecordedAt()), + event.getEvidence()); + versions.add(tombstone); + addSupersedes( + tombstone.getId(), + source, + materializedSourceIds, + relations); + } + } + + for (MemoryFactVersion version : affected) { + for (TimeInterval remaining + : version.getValidTime().subtract( + event.getValidTime())) { + String fragmentId = versionId( + event, + fragmentIndex++); + versions.add(new MemoryFactVersion( + fragmentId, + version.getFact(), + MemoryFactVersionStatus.ACTIVE, + remaining, + TimeInterval.unboundedFrom( + event.getRecordedAt()), + version.getEvidence())); + addSupersedes( + fragmentId, + version, + materializedSourceIds, + relations); + } + } + + addCurrentConflicts(versions, relations); + } + + private static void addCurrentConflicts( + List versions, + Set relations) { + for (int leftIndex = 0; + leftIndex < versions.size(); leftIndex++) { + MemoryFactVersion left = versions.get(leftIndex); + if (!isCurrentActive(left)) { + continue; + } + for (int rightIndex = leftIndex + 1; + rightIndex < versions.size(); rightIndex++) { + MemoryFactVersion right = versions.get(rightIndex); + if (isCurrentActive(right) + && left.getValidTime().overlaps( + right.getValidTime()) + && !factValue(left.getFact()).equals( + factValue(right.getFact()))) { + relations.add(new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + left.getId(), + right.getId())); + } + } + } + } + + private static void addSupersedes( + String replacementId, + List sources, + Set materializedSourceIds, + Set relations) { + for (MemoryFactVersion source : sources) { + addSupersedes( + replacementId, + source, + materializedSourceIds, + relations); + } + } + + private static void addSupersedes( + String replacementId, + MemoryFactVersion source, + Set materializedSourceIds, + Set relations) { + if (materializedSourceIds.contains(source.getId())) { + relations.add(new VersionRelation( + VersionRelationType.SUPERSEDES, + replacementId, + source.getId())); + } + } + + private static boolean closeVersion( + MemoryFactVersion version, + Instant recordedAt, + List versions, + Set relations) { + Instant startedAt = + version.getTransactionTime().getStart(); + if (startedAt.isAfter(recordedAt)) { + throw new IllegalArgumentException( + "Event recorded time precedes current version"); + } + if (startedAt.equals(recordedAt)) { + relations.removeIf(relation -> + references(relation, version.getId())); + return false; + } + versions.add(new MemoryFactVersion( + version.getId(), + version.getFact(), + version.getStatus(), + version.getValidTime(), + new TimeInterval(startedAt, recordedAt), + version.getEvidence())); + return true; + } + + private static boolean references( + VersionRelation relation, + String versionId) { + return relation.getFromVersionId().equals(versionId) + || relation.getToVersionId().equals(versionId); + } + + private static FactValue factValue(MemoryFact fact) { + if (fact.isRelationship()) { + return FactValue.entityReference( + fact.getTarget().get().getId()); + } + return FactValue.literal( + fact.getLiteralValue().get()); + } + + private static TimeInterval span( + TimeInterval left, + TimeInterval right) { + Instant start = left.getStart().isBefore(right.getStart()) + ? left.getStart() : right.getStart(); + Instant end; + if (!left.getEnd().isPresent() + || !right.getEnd().isPresent()) { + end = null; + } else { + Instant leftEnd = left.getEnd().get(); + Instant rightEnd = right.getEnd().get(); + end = leftEnd.isAfter(rightEnd) ? leftEnd : rightEnd; + } + return new TimeInterval(start, end); + } + + private static List sortedUniqueEvidence( + List evidence) { + Collections.sort(evidence, EVIDENCE_ORDER); + List unique = new ArrayList<>(); + for (Evidence item : evidence) { + if (unique.isEmpty() + || !unique.get(unique.size() - 1).equals(item)) { + unique.add(item); + } + } + return unique; + } + + private static String versionId( + NormalizedMemoryEvent event, + int index) { + return event.getEventId() + ":version:" + index; + } + + private static boolean isCurrentActive( + MemoryFactVersion version) { + return isCurrent(version) + && version.getStatus() + == MemoryFactVersionStatus.ACTIVE; + } + + private static void replayAdd( + MemoryEvent event, + List versions) { + for (MemoryFactVersion version : versions) { + if (isCurrent(version) + && version.getFact().getId().equals(event.getFactId()) + && version.getValidTime().overlaps( + event.getValidTime())) { + throw new IllegalArgumentException( + "Overlapping add for fact id: " + + event.getFactId()); + } + } + + versions.add(new MemoryFactVersion( + event.getId() + ":version:0", + event.getFact().get(), + event.getValidTime(), + TimeInterval.unboundedFrom( + event.getTransactionTime()), + event.getEvidence())); + } + + private static void replayChange( + MemoryEvent event, + List versions) { + List affected = new ArrayList<>(); + + for (MemoryFactVersion version : versions) { + if (isCurrent(version) + && version.getFact().getId().equals(event.getFactId()) + && version.getValidTime().overlaps( + event.getValidTime())) { + affected.add(version); + } + } + + Collections.sort(affected, VALID_TIME_ORDER); + + if (!isFullyCovered(event.getValidTime(), affected)) { + throw new IllegalArgumentException( + "Event interval is not fully covered for fact id: " + + event.getFactId()); + } + + versions.removeAll(affected); + + int fragmentIndex = 1; + for (MemoryFactVersion version : affected) { + if (version.getTransactionTime().getStart() + .isBefore(event.getTransactionTime())) { + versions.add(new MemoryFactVersion( + version.getId(), + version.getFact(), + version.getValidTime(), + new TimeInterval( + version.getTransactionTime().getStart(), + event.getTransactionTime()), + version.getEvidence())); + } + + for (TimeInterval remaining : + version.getValidTime().subtract( + event.getValidTime())) { + versions.add(new MemoryFactVersion( + event.getId() + ":version:" + + fragmentIndex++, + version.getFact(), + remaining, + TimeInterval.unboundedFrom( + event.getTransactionTime()), + version.getEvidence())); + } + } + + if (event.getOperation() == MemoryEventOperation.CORRECT) { + versions.add(new MemoryFactVersion( + event.getId() + ":version:0", + event.getFact().get(), + event.getValidTime(), + TimeInterval.unboundedFrom( + event.getTransactionTime()), + event.getEvidence())); + } + } + + private static boolean isFullyCovered( + TimeInterval target, + List coveringVersions) { + List uncovered = new ArrayList<>(); + uncovered.add(target); + + for (MemoryFactVersion version : coveringVersions) { + List remaining = new ArrayList<>(); + + for (TimeInterval interval : uncovered) { + remaining.addAll( + interval.subtract(version.getValidTime())); + } + + uncovered = remaining; + if (uncovered.isEmpty()) { + return true; + } + } + + return false; + } + + private static boolean isCurrent( + MemoryFactVersion version) { + return !version.getTransactionTime() + .getEnd().isPresent(); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/oracle/ReplayMethod.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/oracle/ReplayMethod.java new file mode 100644 index 000000000..bf01a5d19 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/oracle/ReplayMethod.java @@ -0,0 +1,32 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.oracle; + +import java.util.List; +import org.apache.geaflow.ai.temporal.semantics.CanonicalSnapshot; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; + +/** + * Generates a canonical snapshot from normalized memory events. + */ +public interface ReplayMethod { + + CanonicalSnapshot replayToSnapshot(List events); +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/query/BitemporalQuery.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/query/BitemporalQuery.java new file mode 100644 index 000000000..ce88813c0 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/query/BitemporalQuery.java @@ -0,0 +1,65 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.query; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.List; +import java.util.Objects; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; + +/** + * Selects memory fact versions visible at two temporal points. + */ +public final class BitemporalQuery { + + private static final Comparator RESULT_ORDER = + Comparator.comparing( + (MemoryFactVersion version) -> + version.getFact().getId()) + .thenComparing(MemoryFactVersion::getId); + + public List query( + List versions, + Instant validAt, + Instant transactionAt) { + Objects.requireNonNull(versions, "versions"); + Objects.requireNonNull(validAt, "validAt"); + Objects.requireNonNull(transactionAt, "transactionAt"); + + List matches = new ArrayList<>(); + for (MemoryFactVersion version : versions) { + Objects.requireNonNull(version, "version"); + if (version.getStatus() + == MemoryFactVersionStatus.ACTIVE + && version.getValidTime().contains(validAt) + && version.getTransactionTime().contains( + transactionAt)) { + matches.add(version); + } + } + + Collections.sort(matches, RESULT_ORDER); + return Collections.unmodifiableList(matches); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/CanonicalSnapshot.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/CanonicalSnapshot.java new file mode 100644 index 000000000..4b98c92d6 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/CanonicalSnapshot.java @@ -0,0 +1,239 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.HashMap; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Set; +import java.util.TreeMap; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.VersionRelation; + +/** + * An immutable canonical representation of temporal state and its input events. + */ +public final class CanonicalSnapshot { + + private static final Comparator EVENT_ORDER = + Comparator.comparing(NormalizedMemoryEvent::getRecordedAt) + .thenComparing(NormalizedMemoryEvent::getEventId); + + private static final Comparator EVIDENCE_ORDER = + Comparator.comparing(Evidence::getId) + .thenComparing(evidence -> evidence.getSource().getId()) + .thenComparing(evidence -> evidence.getSource().getName()) + .thenComparing(Evidence::getContent); + + private final TemporalState state; + private final List events; + private final Map generatingEventIds; + + public CanonicalSnapshot( + TemporalState state, + List events, + Map generatingEventIds) { + this.state = canonicalizeState( + Objects.requireNonNull(state, "state")); + this.events = canonicalizeEvents(events); + this.generatingEventIds = canonicalizeGeneratingEventIds( + generatingEventIds, + this.state, + this.events); + } + + public TemporalState getState() { + return state; + } + + public List getEvents() { + return events; + } + + public Map getGeneratingEventIds() { + return generatingEventIds; + } + + private static TemporalState canonicalizeState( + TemporalState state) { + Map> versionsByKey = + new TreeMap<>(); + Set versionIds = new HashSet<>(); + for (List versions + : state.getVersionsByFactKey().values()) { + for (MemoryFactVersion version : versions) { + if (!versionIds.add(version.getId())) { + throw new IllegalArgumentException( + "Duplicate version id: " + version.getId()); + } + } + } + for (Map.Entry> entry + : state.getVersionsByFactKey().entrySet()) { + FactKey factKey = Objects.requireNonNull( + entry.getKey(), "factKey"); + List versions = Objects.requireNonNull( + entry.getValue(), "versions"); + if (versions.isEmpty()) { + continue; + } + List canonicalVersions = + new ArrayList<>(); + for (MemoryFactVersion version : versions) { + Objects.requireNonNull(version, "version"); + validateFactKey(factKey, version.getFact()); + canonicalVersions.add(canonicalizeVersion(version)); + } + versionsByKey.put(factKey, canonicalVersions); + } + for (VersionRelation relation : state.getRelations()) { + Objects.requireNonNull(relation, "relation"); + if (!versionIds.contains(relation.getFromVersionId()) + || !versionIds.contains(relation.getToVersionId())) { + throw new IllegalArgumentException( + "Version relation references an unknown version"); + } + } + return new TemporalState(versionsByKey, state.getRelations()); + } + + private static MemoryFactVersion canonicalizeVersion( + MemoryFactVersion version) { + List evidence = new ArrayList<>(version.getEvidence()); + for (Evidence item : evidence) { + Objects.requireNonNull(item, "evidence"); + } + Collections.sort(evidence, EVIDENCE_ORDER); + return new MemoryFactVersion( + version.getId(), + version.getFact(), + version.getStatus(), + version.getValidTime(), + version.getTransactionTime(), + evidence); + } + + private static void validateFactKey( + FactKey factKey, + MemoryFact fact) { + if (!factKey.getSubjectId().equals(fact.getSubject().getId()) + || !factKey.getPredicate().equals(fact.getPredicate())) { + throw new IllegalArgumentException( + "Fact key does not match version fact"); + } + } + + private static List canonicalizeEvents( + List events) { + List ordered = new ArrayList<>( + Objects.requireNonNull(events, "events")); + for (NormalizedMemoryEvent event : ordered) { + Objects.requireNonNull(event, "event"); + } + Collections.sort(ordered, EVENT_ORDER); + + Map payloadHashes = new HashMap<>(); + List unique = new ArrayList<>(); + for (NormalizedMemoryEvent event : ordered) { + String previousHash = payloadHashes.get(event.getEventId()); + if (previousHash == null) { + payloadHashes.put( + event.getEventId(), event.getPayloadHash()); + unique.add(event); + } else if (!previousHash.equals(event.getPayloadHash())) { + throw new IllegalArgumentException( + "Event id reused with a different payload: " + + event.getEventId()); + } + } + return Collections.unmodifiableList(unique); + } + + private static Map canonicalizeGeneratingEventIds( + Map generatingEventIds, + TemporalState state, + List events) { + Map versionKeys = new HashMap<>(); + for (Map.Entry> entry + : state.getVersionsByFactKey().entrySet()) { + for (MemoryFactVersion version : entry.getValue()) { + versionKeys.put(version.getId(), entry.getKey()); + } + } + Map eventsById = new HashMap<>(); + for (NormalizedMemoryEvent event : events) { + eventsById.put(event.getEventId(), event); + } + + Map ordered = new TreeMap<>(); + for (Map.Entry entry : Objects.requireNonNull( + generatingEventIds, "generatingEventIds").entrySet()) { + String versionId = Objects.requireNonNull( + entry.getKey(), "versionId"); + String eventId = Objects.requireNonNull( + entry.getValue(), "eventId"); + ordered.put(versionId, eventId); + } + if (!ordered.keySet().equals(versionKeys.keySet())) { + throw new IllegalArgumentException( + "Generating event ids must match version ids"); + } + for (Map.Entry entry : ordered.entrySet()) { + NormalizedMemoryEvent event = eventsById.get(entry.getValue()); + if (event == null) { + throw new IllegalArgumentException( + "Generating event is unknown: " + entry.getValue()); + } + if (!versionKeys.get(entry.getKey()).equals( + event.getFactKey())) { + throw new IllegalArgumentException( + "Generating event fact key does not match version"); + } + } + return Collections.unmodifiableMap(ordered); + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof CanonicalSnapshot)) { + return false; + } + CanonicalSnapshot that = (CanonicalSnapshot) object; + return state.equals(that.state) + && events.equals(that.events) + && generatingEventIds.equals(that.generatingEventIds); + } + + @Override + public int hashCode() { + return Objects.hash(state, events, generatingEventIds); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/DiffReport.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/DiffReport.java new file mode 100644 index 000000000..eb6c2ca14 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/DiffReport.java @@ -0,0 +1,107 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.util.Objects; +import java.util.Optional; +import org.apache.geaflow.ai.temporal.model.FactKey; + +/** + * The first difference found between two canonical snapshots. + */ +public final class DiffReport { + + private final boolean equivalent; + private final String fieldPath; + private final String expectedValue; + private final String actualValue; + private final FactKey factKey; + private final String eventId; + + private DiffReport(boolean equivalent, String fieldPath, String expectedValue, + String actualValue, FactKey factKey, String eventId) { + this.equivalent = equivalent; + this.fieldPath = fieldPath; + this.expectedValue = expectedValue; + this.actualValue = actualValue; + this.factKey = factKey; + this.eventId = eventId; + } + + public static DiffReport equivalent() { + return new DiffReport(true, null, null, null, null, null); + } + + public static DiffReport difference(String fieldPath, String expectedValue, + String actualValue, FactKey factKey, + String eventId) { + Objects.requireNonNull(fieldPath, "fieldPath"); + if (fieldPath.trim().isEmpty()) { + throw new IllegalArgumentException("Difference field path must not be blank"); + } + return new DiffReport(false, fieldPath, expectedValue, actualValue, factKey, eventId); + } + + public boolean isEquivalent() { + return equivalent; + } + + public Optional getFieldPath() { + return Optional.ofNullable(fieldPath); + } + + public Optional getExpectedValue() { + return Optional.ofNullable(expectedValue); + } + + public Optional getActualValue() { + return Optional.ofNullable(actualValue); + } + + public Optional getFactKey() { + return Optional.ofNullable(factKey); + } + + public Optional getEventId() { + return Optional.ofNullable(eventId); + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof DiffReport)) { + return false; + } + DiffReport that = (DiffReport) object; + return equivalent == that.equivalent + && Objects.equals(fieldPath, that.fieldPath) + && Objects.equals(expectedValue, that.expectedValue) + && Objects.equals(actualValue, that.actualValue) + && Objects.equals(factKey, that.factKey) + && Objects.equals(eventId, that.eventId); + } + + @Override + public int hashCode() { + return Objects.hash(equivalent, fieldPath, expectedValue, actualValue, factKey, eventId); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventLedger.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventLedger.java new file mode 100644 index 000000000..3d021549f --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventLedger.java @@ -0,0 +1,65 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; + +/** + * Tracks the committed payload hash for each normalized event id. + */ +public final class EventLedger { + + private final Map payloadHashes = new HashMap<>(); + + public EventLedgerDecision check(NormalizedMemoryEvent event) { + Objects.requireNonNull(event, "event"); + String existing = payloadHashes.get(event.getEventId()); + if (existing == null) { + return EventLedgerDecision.ACCEPTED; + } + if (existing.equals(event.getPayloadHash())) { + return EventLedgerDecision.DUPLICATE_NOOP; + } + return EventLedgerDecision.REJECT_EVENT_ID_REUSE; + } + + public EventLedgerDecision commit(NormalizedMemoryEvent event) { + EventLedgerDecision decision = check(event); + if (decision == EventLedgerDecision.REJECT_EVENT_ID_REUSE) { + throw new IllegalArgumentException( + "Event id reused with a different payload: " + + event.getEventId()); + } + if (decision == EventLedgerDecision.ACCEPTED) { + payloadHashes.put( + event.getEventId(), + event.getPayloadHash()); + } + return decision; + } + + public Optional getPayloadHash(String eventId) { + return Optional.ofNullable(payloadHashes.get( + Objects.requireNonNull(eventId, "eventId"))); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventLedgerDecision.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventLedgerDecision.java new file mode 100644 index 000000000..22a387070 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventLedgerDecision.java @@ -0,0 +1,30 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +/** + * Result of checking a normalized event against the event ledger. + */ +public enum EventLedgerDecision { + + ACCEPTED, + DUPLICATE_NOOP, + REJECT_EVENT_ID_REUSE +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventNormalizer.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventNormalizer.java new file mode 100644 index 000000000..52623b321 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/EventNormalizer.java @@ -0,0 +1,347 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; +import java.security.MessageDigest; +import java.security.NoSuchAlgorithmException; +import java.text.Normalizer; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.List; +import java.util.Objects; +import java.util.Optional; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; + +/** + * Canonicalizes temporal events before replay or ledger checks. + */ +public final class EventNormalizer { + + private static final char[] HEX = "0123456789abcdef".toCharArray(); + + private static final Comparator EVIDENCE_ORDER = + Comparator.comparing(Evidence::getId) + .thenComparing(evidence -> evidence.getSource().getId()) + .thenComparing(evidence -> evidence.getSource().getName()) + .thenComparing(Evidence::getContent); + + public NormalizedMemoryEvent normalize( + MemoryEvent event, + FactKey factKey) { + Objects.requireNonNull(event, "event"); + Objects.requireNonNull(factKey, "factKey"); + + FactKey normalizedKey = new FactKey( + normalizeText(factKey.getSubjectId()), + normalizeText(factKey.getPredicate()), + normalizeText(factKey.getScope())); + String eventId = normalizeText(event.getId()); + String factId = normalizeText(event.getFactId()); + TimeInterval validTime = normalizeInterval( + event.getValidTime()); + Instant recordedAt = normalizeTime( + event.getTransactionTime()); + List evidence = normalizeEvidence( + event.getEvidence()); + MemoryFact fact = normalizeFact(event, factId); + + validateFactKey(event.getOperation(), fact, normalizedKey); + MemoryEvent normalizedEvent = createEvent( + eventId, + event.getOperation(), + factId, + fact, + validTime, + recordedAt, + evidence); + FactValue factValue = createFactValue(fact); + String payloadHash = payloadHash( + normalizedEvent, + normalizedKey, + factValue); + + return new NormalizedMemoryEvent( + normalizedEvent, + normalizedKey, + factValue, + payloadHash); + } + + private static MemoryFact normalizeFact( + MemoryEvent event, + String factId) { + Optional optionalFact = event.getFact(); + if (!optionalFact.isPresent()) { + if (event.getOperation() != MemoryEventOperation.RETRACT) { + throw new IllegalArgumentException( + "Event operation requires a fact"); + } + return null; + } + + MemoryFact fact = optionalFact.get(); + MemoryEntity subject = normalizeEntity(fact.getSubject()); + String predicate = normalizeText(fact.getPredicate()); + if (fact.isRelationship()) { + return MemoryFact.relationship( + factId, + subject, + predicate, + normalizeEntity(fact.getTarget().get())); + } + return MemoryFact.attribute( + factId, + subject, + predicate, + normalizeText(fact.getLiteralValue().get())); + } + + private static MemoryEntity normalizeEntity( + MemoryEntity entity) { + return new MemoryEntity( + normalizeText(entity.getId()), + normalizeText(entity.getLabel())); + } + + private static List normalizeEvidence( + List evidence) { + List normalized = new ArrayList<>(); + for (Evidence item : evidence) { + Source source = item.getSource(); + normalized.add(new Evidence( + normalizeText(item.getId()), + new Source( + normalizeText(source.getId()), + normalizeText(source.getName())), + normalizeText(item.getContent()))); + } + Collections.sort(normalized, EVIDENCE_ORDER); + return normalized; + } + + private static TimeInterval normalizeInterval( + TimeInterval interval) { + Instant end = interval.getEnd().isPresent() + ? normalizeTime(interval.getEnd().get()) : null; + return new TimeInterval( + normalizeTime(interval.getStart()), + end); + } + + private static Instant normalizeTime(Instant time) { + return time; + } + + private static String normalizeText(String value) { + return Normalizer.normalize(value, Normalizer.Form.NFC); + } + + private static void validateFactKey( + MemoryEventOperation operation, + MemoryFact fact, + FactKey factKey) { + if (operation == MemoryEventOperation.RETRACT) { + if (fact != null) { + throw new IllegalArgumentException( + "Retract event must not contain a fact"); + } + return; + } + if (fact == null + || !fact.getSubject().getId().equals( + factKey.getSubjectId()) + || !fact.getPredicate().equals( + factKey.getPredicate())) { + throw new IllegalArgumentException( + "Fact key does not match event fact"); + } + } + + private static MemoryEvent createEvent( + String eventId, + MemoryEventOperation operation, + String factId, + MemoryFact fact, + TimeInterval validTime, + Instant recordedAt, + List evidence) { + if (operation == MemoryEventOperation.ADD) { + return MemoryEvent.add( + eventId, + fact, + validTime, + recordedAt, + evidence); + } + if (operation == MemoryEventOperation.CORRECT) { + return MemoryEvent.correct( + eventId, + fact, + validTime, + recordedAt, + evidence); + } + if (operation == MemoryEventOperation.RETRACT) { + return MemoryEvent.retract( + eventId, + factId, + validTime, + recordedAt, + evidence); + } + throw new UnsupportedOperationException( + "Unsupported memory event operation: " + operation); + } + + private static FactValue createFactValue(MemoryFact fact) { + if (fact == null) { + return null; + } + if (fact.isRelationship()) { + return FactValue.entityReference( + fact.getTarget().get().getId()); + } + return FactValue.literal( + fact.getLiteralValue().get()); + } + + private static String payloadHash( + MemoryEvent event, + FactKey factKey, + FactValue factValue) { + MessageDigest digest = sha256(); + updateText(digest, event.getOperation().name()); + updateText(digest, event.getFactId()); + updateText(digest, factKey.getSubjectId()); + updateText(digest, factKey.getPredicate()); + updateText(digest, factKey.getScope()); + updateText( + digest, + factValue == null ? null : factValue.getKind().name()); + updateText( + digest, + factValue == null ? null : factValue.getValue()); + + MemoryFact fact = event.getFact().orElse(null); + updateText( + digest, + fact == null ? null : fact.getSubject().getLabel()); + updateText( + digest, + fact != null && fact.isRelationship() + ? fact.getTarget().get().getLabel() : null); + updateTime( + digest, + event.getValidTime().getStart()); + updateOptionalTime( + digest, + event.getValidTime().getEnd()); + updateTime( + digest, + event.getTransactionTime()); + updateInt(digest, event.getEvidence().size()); + for (Evidence evidence : event.getEvidence()) { + updateText(digest, evidence.getId()); + updateText(digest, evidence.getSource().getId()); + updateText(digest, evidence.getSource().getName()); + updateText(digest, evidence.getContent()); + } + return toHex(digest.digest()); + } + + private static MessageDigest sha256() { + try { + return MessageDigest.getInstance("SHA-256"); + } catch (NoSuchAlgorithmException exception) { + throw new IllegalStateException( + "SHA-256 is unavailable", + exception); + } + } + + private static void updateText( + MessageDigest digest, + String value) { + if (value == null) { + digest.update((byte) 0); + return; + } + digest.update((byte) 1); + byte[] bytes = value.getBytes(StandardCharsets.UTF_8); + updateInt(digest, bytes.length); + digest.update(bytes); + } + + private static void updateOptionalTime( + MessageDigest digest, + Optional time) { + if (time.isPresent()) { + digest.update((byte) 1); + updateTime(digest, time.get()); + } else { + digest.update((byte) 0); + } + } + + private static void updateTime( + MessageDigest digest, + Instant time) { + updateLong(digest, time.getEpochSecond()); + updateInt(digest, time.getNano()); + } + + private static void updateInt( + MessageDigest digest, + int value) { + digest.update(ByteBuffer.allocate(Integer.BYTES) + .putInt(value) + .array()); + } + + private static void updateLong( + MessageDigest digest, + long value) { + digest.update(ByteBuffer.allocate(Long.BYTES) + .putLong(value) + .array()); + } + + private static String toHex(byte[] bytes) { + char[] characters = new char[bytes.length * 2]; + for (int index = 0; index < bytes.length; index++) { + int value = bytes[index] & 0xff; + characters[index * 2] = HEX[value >>> 4]; + characters[index * 2 + 1] = HEX[value & 0x0f]; + } + return new String(characters); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/NormalizedMemoryEvent.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/NormalizedMemoryEvent.java new file mode 100644 index 000000000..5bd294647 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/NormalizedMemoryEvent.java @@ -0,0 +1,116 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.time.Instant; +import java.util.List; +import java.util.Objects; +import java.util.Optional; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.TimeInterval; + +/** + * An immutable event after canonical input normalization. + */ +public final class NormalizedMemoryEvent { + + private final MemoryEvent event; + private final FactKey factKey; + private final FactValue factValue; + private final String payloadHash; + + NormalizedMemoryEvent( + MemoryEvent event, + FactKey factKey, + FactValue factValue, + String payloadHash) { + this.event = Objects.requireNonNull(event, "event"); + this.factKey = Objects.requireNonNull(factKey, "factKey"); + this.factValue = factValue; + this.payloadHash = Objects.requireNonNull( + payloadHash, + "payloadHash"); + } + + public MemoryEvent getEvent() { + return event; + } + + public String getEventId() { + return event.getId(); + } + + public MemoryEventOperation getOperation() { + return event.getOperation(); + } + + public String getFactId() { + return event.getFactId(); + } + + public FactKey getFactKey() { + return factKey; + } + + public Optional getFactValue() { + return Optional.ofNullable(factValue); + } + + public TimeInterval getValidTime() { + return event.getValidTime(); + } + + public Instant getRecordedAt() { + return event.getTransactionTime(); + } + + public List getEvidence() { + return event.getEvidence(); + } + + public String getPayloadHash() { + return payloadHash; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof NormalizedMemoryEvent)) { + return false; + } + NormalizedMemoryEvent that = + (NormalizedMemoryEvent) object; + return event.equals(that.event) + && factKey.equals(that.factKey) + && Objects.equals(factValue, that.factValue) + && payloadHash.equals(that.payloadHash); + } + + @Override + public int hashCode() { + return Objects.hash(event, factKey, factValue, payloadHash); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/SnapshotComparator.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/SnapshotComparator.java new file mode 100644 index 000000000..10db72103 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/SnapshotComparator.java @@ -0,0 +1,632 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Set; +import java.util.TreeMap; +import java.util.TreeSet; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; + +/** + * Finds the first stable, field-level difference between two canonical snapshots. + */ +public final class SnapshotComparator { + + public DiffReport compare( + CanonicalSnapshot expected, + CanonicalSnapshot actual) { + Objects.requireNonNull(expected, "expected"); + Objects.requireNonNull(actual, "actual"); + + DiffReport difference = compareEvents(expected, actual); + if (difference != null) { + return difference; + } + difference = compareVersions(expected, actual); + if (difference != null) { + return difference; + } + difference = compareGeneratingEventIds(expected, actual); + if (difference != null) { + return difference; + } + difference = compareRelations(expected, actual); + return difference == null ? DiffReport.equivalent() : difference; + } + + private static DiffReport compareEvents( + CanonicalSnapshot expected, + CanonicalSnapshot actual) { + Map expectedEvents = eventsById(expected); + Map actualEvents = eventsById(actual); + Set eventIds = keys(expectedEvents, actualEvents); + for (String eventId : eventIds) { + NormalizedMemoryEvent expectedEvent = expectedEvents.get(eventId); + NormalizedMemoryEvent actualEvent = actualEvents.get(eventId); + NormalizedMemoryEvent contextEvent = expectedEvent == null + ? actualEvent : expectedEvent; + FactKey contextKey = contextEvent.getFactKey(); + String prefix = "events[" + eventId + "]"; + DiffReport difference = comparePresence( + prefix, expectedEvent, actualEvent, contextKey, eventId); + if (difference != null) { + return difference; + } + difference = compareFactKey( + prefix + ".factKey", + expectedEvent.getFactKey(), + actualEvent.getFactKey(), + contextKey, + eventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".operation", + expectedEvent.getOperation().name(), + actualEvent.getOperation().name(), + contextKey, + eventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".factId", + expectedEvent.getFactId(), + actualEvent.getFactId(), + contextKey, + eventId); + if (difference != null) { + return difference; + } + difference = compareOptionalFact( + prefix + ".fact", + expectedEvent.getEvent().getFact().orElse(null), + actualEvent.getEvent().getFact().orElse(null), + contextKey, + eventId); + if (difference != null) { + return difference; + } + difference = compareInterval( + prefix + ".validTime", + expectedEvent.getValidTime(), + actualEvent.getValidTime(), + contextKey, + eventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".recordedAt", + time(expectedEvent.getRecordedAt()), + time(actualEvent.getRecordedAt()), + contextKey, + eventId); + if (difference != null) { + return difference; + } + difference = compareEvidence( + prefix + ".evidence", + expectedEvent.getEvidence(), + actualEvent.getEvidence(), + contextKey, + eventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".payloadHash", + expectedEvent.getPayloadHash(), + actualEvent.getPayloadHash(), + contextKey, + eventId); + if (difference != null) { + return difference; + } + } + return null; + } + + private static DiffReport compareVersions( + CanonicalSnapshot expected, + CanonicalSnapshot actual) { + Map expectedVersions = versionsById(expected); + Map actualVersions = versionsById(actual); + Set versionIds = keys(expectedVersions, actualVersions); + for (String versionId : versionIds) { + VersionEntry expectedEntry = expectedVersions.get(versionId); + VersionEntry actualEntry = actualVersions.get(versionId); + VersionEntry contextEntry = expectedEntry == null + ? actualEntry : expectedEntry; + CanonicalSnapshot contextSnapshot = expectedEntry == null + ? actual : expected; + FactKey contextKey = contextEntry.factKey; + String contextEventId = contextSnapshot.getGeneratingEventIds().get(versionId); + String prefix = "versions[" + versionId + "]"; + DiffReport difference = comparePresence( + prefix, expectedEntry, actualEntry, contextKey, contextEventId); + if (difference != null) { + return difference; + } + difference = compareFactKey( + prefix + ".factKey", + expectedEntry.factKey, + actualEntry.factKey, + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + MemoryFactVersion expectedVersion = expectedEntry.version; + MemoryFactVersion actualVersion = actualEntry.version; + difference = compareFact( + prefix + ".fact", + expectedVersion.getFact(), + actualVersion.getFact(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".status", + expectedVersion.getStatus().name(), + actualVersion.getStatus().name(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = compareInterval( + prefix + ".validTime", + expectedVersion.getValidTime(), + actualVersion.getValidTime(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = compareInterval( + prefix + ".transactionTime", + expectedVersion.getTransactionTime(), + actualVersion.getTransactionTime(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = compareEvidence( + prefix + ".evidence", + expectedVersion.getEvidence(), + actualVersion.getEvidence(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + } + return null; + } + + private static DiffReport compareGeneratingEventIds( + CanonicalSnapshot expected, + CanonicalSnapshot actual) { + Map expectedIds = expected.getGeneratingEventIds(); + Map actualIds = actual.getGeneratingEventIds(); + Map expectedVersions = versionsById(expected); + Map actualVersions = versionsById(actual); + for (String versionId : keys(expectedIds, actualIds)) { + String expectedEventId = expectedIds.get(versionId); + String actualEventId = actualIds.get(versionId); + VersionEntry contextEntry = expectedVersions.get(versionId); + if (contextEntry == null) { + contextEntry = actualVersions.get(versionId); + } + FactKey contextKey = contextEntry == null + ? null : contextEntry.factKey; + DiffReport difference = difference( + "generatingEventIds[" + versionId + "]", + expectedEventId, + actualEventId, + contextKey, + expectedEventId); + if (difference != null) { + return difference; + } + } + return null; + } + + private static DiffReport compareRelations( + CanonicalSnapshot expected, + CanonicalSnapshot actual) { + List expectedRelations = new ArrayList<>( + expected.getState().getRelations()); + List actualRelations = new ArrayList<>( + actual.getState().getRelations()); + Collections.sort(expectedRelations); + Collections.sort(actualRelations); + Map expectedVersions = versionsById(expected); + Map actualVersions = versionsById(actual); + int relationCount = Math.max(expectedRelations.size(), actualRelations.size()); + for (int index = 0; index < relationCount; index++) { + VersionRelation expectedRelation = index < expectedRelations.size() + ? expectedRelations.get(index) : null; + VersionRelation actualRelation = index < actualRelations.size() + ? actualRelations.get(index) : null; + VersionRelation contextRelation = expectedRelation == null + ? actualRelation : expectedRelation; + CanonicalSnapshot contextSnapshot = expectedRelation == null + ? actual : expected; + Map contextVersions = expectedRelation == null + ? actualVersions : expectedVersions; + VersionEntry contextEntry = contextVersions.get( + contextRelation.getFromVersionId()); + FactKey contextKey = contextEntry == null + ? null : contextEntry.factKey; + String contextEventId = contextSnapshot.getGeneratingEventIds().get( + contextRelation.getFromVersionId()); + String prefix = "relations[" + index + "]"; + DiffReport difference = comparePresence( + prefix, + expectedRelation, + actualRelation, + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".type", + expectedRelation.getType().name(), + actualRelation.getType().name(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".fromVersionId", + expectedRelation.getFromVersionId(), + actualRelation.getFromVersionId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".toVersionId", + expectedRelation.getToVersionId(), + actualRelation.getToVersionId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + } + return null; + } + + private static DiffReport compareFactKey( + String prefix, + FactKey expected, + FactKey actual, + FactKey contextKey, + String contextEventId) { + DiffReport difference = difference( + prefix + ".subjectId", + expected.getSubjectId(), + actual.getSubjectId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".predicate", + expected.getPredicate(), + actual.getPredicate(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + return difference( + prefix + ".scope", + expected.getScope(), + actual.getScope(), + contextKey, + contextEventId); + } + + private static DiffReport compareOptionalFact( + String prefix, + MemoryFact expected, + MemoryFact actual, + FactKey contextKey, + String contextEventId) { + DiffReport difference = comparePresence( + prefix, expected, actual, contextKey, contextEventId); + if (difference != null || expected == null) { + return difference; + } + return compareFact( + prefix, expected, actual, contextKey, contextEventId); + } + + private static DiffReport compareFact( + String prefix, + MemoryFact expected, + MemoryFact actual, + FactKey contextKey, + String contextEventId) { + DiffReport difference = difference( + prefix + ".id", + expected.getId(), + actual.getId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".subject.id", + expected.getSubject().getId(), + actual.getSubject().getId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".subject.label", + expected.getSubject().getLabel(), + actual.getSubject().getLabel(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".predicate", + expected.getPredicate(), + actual.getPredicate(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + prefix + ".kind", + factKind(expected), + factKind(actual), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + if (!expected.isRelationship()) { + return difference( + prefix + ".literalValue", + expected.getLiteralValue().orElse(null), + actual.getLiteralValue().orElse(null), + contextKey, + contextEventId); + } + difference = difference( + prefix + ".target.id", + expected.getTarget().get().getId(), + actual.getTarget().get().getId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + return difference( + prefix + ".target.label", + expected.getTarget().get().getLabel(), + actual.getTarget().get().getLabel(), + contextKey, + contextEventId); + } + + private static DiffReport compareInterval( + String prefix, + TimeInterval expected, + TimeInterval actual, + FactKey contextKey, + String contextEventId) { + DiffReport difference = difference( + prefix + ".start", + time(expected.getStart()), + time(actual.getStart()), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + return difference( + prefix + ".end", + end(expected), + end(actual), + contextKey, + contextEventId); + } + + private static DiffReport compareEvidence( + String prefix, + List expected, + List actual, + FactKey contextKey, + String contextEventId) { + int evidenceCount = Math.max(expected.size(), actual.size()); + for (int index = 0; index < evidenceCount; index++) { + Evidence expectedEvidence = index < expected.size() + ? expected.get(index) : null; + Evidence actualEvidence = index < actual.size() + ? actual.get(index) : null; + String itemPrefix = prefix + "[" + index + "]"; + DiffReport difference = comparePresence( + itemPrefix, + expectedEvidence, + actualEvidence, + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + itemPrefix + ".id", + expectedEvidence.getId(), + actualEvidence.getId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + itemPrefix + ".source.id", + expectedEvidence.getSource().getId(), + actualEvidence.getSource().getId(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + itemPrefix + ".source.name", + expectedEvidence.getSource().getName(), + actualEvidence.getSource().getName(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + difference = difference( + itemPrefix + ".content", + expectedEvidence.getContent(), + actualEvidence.getContent(), + contextKey, + contextEventId); + if (difference != null) { + return difference; + } + } + return null; + } + + private static DiffReport comparePresence( + String path, + Object expected, + Object actual, + FactKey contextKey, + String contextEventId) { + return difference( + path, + expected == null ? null : "present", + actual == null ? null : "present", + contextKey, + contextEventId); + } + + private static DiffReport difference( + String path, + String expected, + String actual, + FactKey contextKey, + String contextEventId) { + return Objects.equals(expected, actual) + ? null + : DiffReport.difference( + path, expected, actual, contextKey, contextEventId); + } + + private static String factKind(MemoryFact fact) { + return fact.isRelationship() + ? FactValue.Kind.ENTITY_REF.name() + : FactValue.Kind.LITERAL.name(); + } + + private static String time(Instant instant) { + return instant.toString(); + } + + private static String end(TimeInterval interval) { + return interval.getEnd().isPresent() + ? time(interval.getEnd().get()) : "infinity"; + } + + private static Map eventsById( + CanonicalSnapshot snapshot) { + Map events = new TreeMap<>(); + for (NormalizedMemoryEvent event : snapshot.getEvents()) { + events.put(event.getEventId(), event); + } + return events; + } + + private static Map versionsById( + CanonicalSnapshot snapshot) { + Map versions = new TreeMap<>(); + for (Map.Entry> entry + : snapshot.getState().getVersionsByFactKey().entrySet()) { + for (MemoryFactVersion version : entry.getValue()) { + versions.put( + version.getId(), + new VersionEntry(entry.getKey(), version)); + } + } + return versions; + } + + private static Set keys( + Map expected, + Map actual) { + Set keys = new TreeSet<>(expected.keySet()); + keys.addAll(actual.keySet()); + return keys; + } + + private static final class VersionEntry { + + private final FactKey factKey; + private final MemoryFactVersion version; + + private VersionEntry( + FactKey factKey, + MemoryFactVersion version) { + this.factKey = factKey; + this.version = version; + } + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/TemporalState.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/TemporalState.java new file mode 100644 index 000000000..3cbc2d9a6 --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/semantics/TemporalState.java @@ -0,0 +1,124 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.TreeMap; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.VersionRelation; + +/** + * An immutable, deterministically ordered temporal state snapshot. + */ +public final class TemporalState { + + private static final Comparator VERSION_ORDER = + Comparator.comparing( + (MemoryFactVersion version) -> + version.getTransactionTime().getStart()) + .thenComparing( + version -> version.getValidTime().getStart()) + .thenComparing(MemoryFactVersion::getId); + + private final Map> + versionsByFactKey; + private final List versions; + private final List relations; + + public TemporalState( + Map> versionsByFactKey, + List relations) { + Objects.requireNonNull(versionsByFactKey, "versionsByFactKey"); + + Map> ordered = + new TreeMap<>(); + for (Map.Entry> entry + : versionsByFactKey.entrySet()) { + FactKey key = Objects.requireNonNull( + entry.getKey(), + "factKey"); + List versionsForKey = + new ArrayList<>(Objects.requireNonNull( + entry.getValue(), + "versionsForKey")); + for (MemoryFactVersion version : versionsForKey) { + Objects.requireNonNull(version, "version"); + } + Collections.sort(versionsForKey, VERSION_ORDER); + List immutableVersions = + Collections.unmodifiableList(versionsForKey); + ordered.put(key, immutableVersions); + } + List flattened = new ArrayList<>(); + for (List versionsForKey + : ordered.values()) { + flattened.addAll(versionsForKey); + } + this.versionsByFactKey = Collections.unmodifiableMap(ordered); + this.versions = Collections.unmodifiableList(flattened); + + List orderedRelations = + new ArrayList<>(Objects.requireNonNull( + relations, + "relations")); + for (VersionRelation relation : orderedRelations) { + Objects.requireNonNull(relation, "relation"); + } + Collections.sort(orderedRelations); + this.relations = Collections.unmodifiableList(orderedRelations); + } + + public Map> + getVersionsByFactKey() { + return versionsByFactKey; + } + + public List getVersions() { + return versions; + } + + public List getRelations() { + return relations; + } + + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof TemporalState)) { + return false; + } + TemporalState that = (TemporalState) object; + return versionsByFactKey.equals(that.versionsByFactKey) + && relations.equals(that.relations); + } + + @Override + public int hashCode() { + return Objects.hash(versionsByFactKey, relations); + } +} diff --git a/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/udga/TemporalUdgaProbe.java b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/udga/TemporalUdgaProbe.java new file mode 100644 index 000000000..35472c60d --- /dev/null +++ b/geaflow-ai/src/main/java/org/apache/geaflow/ai/temporal/udga/TemporalUdgaProbe.java @@ -0,0 +1,151 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.udga; + +import java.util.Iterator; +import java.util.List; +import java.util.Optional; +import org.apache.geaflow.common.type.primitive.BooleanType; +import org.apache.geaflow.common.type.primitive.IntegerType; +import org.apache.geaflow.dsl.common.algo.AlgorithmRuntimeContext; +import org.apache.geaflow.dsl.common.algo.AlgorithmUserFunction; +import org.apache.geaflow.dsl.common.algo.IncrementalAlgorithmUserFunction; +import org.apache.geaflow.dsl.common.data.Row; +import org.apache.geaflow.dsl.common.data.RowEdge; +import org.apache.geaflow.dsl.common.data.RowVertex; +import org.apache.geaflow.dsl.common.data.impl.ObjectRow; +import org.apache.geaflow.dsl.common.types.GraphSchema; +import org.apache.geaflow.dsl.common.types.StructType; +import org.apache.geaflow.dsl.common.types.TableField; +import org.apache.geaflow.model.graph.edge.EdgeDirection; + +public class TemporalUdgaProbe implements + AlgorithmUserFunction, + IncrementalAlgorithmUserFunction { + + private AlgorithmRuntimeContext context; + + @Override + public void init( + AlgorithmRuntimeContext context, + Object[] params) { + this.context = context; + } + + @Override + public void process( + RowVertex vertex, + Optional updatedValues, + Iterator messages) { + if (context.getCurrentIterationId() == 1L) { + processDynamicEdges(updatedValues); + } else { + processMessages(updatedValues, messages); + } + } + + private void processDynamicEdges(Optional updatedValues) { + int batchCount = updatedValues + .map(value -> integerField(value, 0)) + .orElse(0) + 1; + int receivedMessageCount = updatedValues + .map(value -> integerField(value, 3)) + .orElse(0); + List dynamicEdges = + context.loadDynamicEdges(EdgeDirection.OUT); + + context.updateVertexValue(ObjectRow.create( + batchCount, + updatedValues.isPresent(), + dynamicEdges.size(), + receivedMessageCount)); + for (RowEdge edge : dynamicEdges) { + context.sendMessage(edge.getTargetId(), 1); + } + } + + private void processMessages( + Optional updatedValues, + Iterator messages) { + int receivedNow = 0; + while (messages.hasNext()) { + receivedNow += messages.next(); + } + if (receivedNow == 0 || !updatedValues.isPresent()) { + return; + } + + Row value = updatedValues.get(); + context.updateVertexValue(ObjectRow.create( + integerField(value, 0), + booleanField(value, 1), + integerField(value, 2), + integerField(value, 3) + receivedNow)); + } + + @Override + public void finish( + RowVertex vertex, + Optional updatedValues) { + if (!updatedValues.isPresent()) { + return; + } + Row value = updatedValues.get(); + context.take(ObjectRow.create( + vertex.getId(), + integerField(value, 0), + booleanField(value, 1), + integerField(value, 2), + integerField(value, 3))); + } + + @Override + public StructType getOutputType(GraphSchema graphSchema) { + return new StructType( + new TableField( + "vertex_id", + graphSchema.getIdType(), + false), + new TableField( + "batch_count", + IntegerType.INSTANCE, + false), + new TableField( + "had_previous_value", + BooleanType.INSTANCE, + false), + new TableField( + "dynamic_edge_count", + IntegerType.INSTANCE, + false), + new TableField( + "received_message_count", + IntegerType.INSTANCE, + false)); + } + + private static int integerField(Row value, int index) { + return (int) value.getField(index, IntegerType.INSTANCE); + } + + private static boolean booleanField(Row value, int index) { + return (boolean) value.getField(index, BooleanType.INSTANCE); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/adapter/MemoryGraphAdapterTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/adapter/MemoryGraphAdapterTest.java new file mode 100644 index 000000000..91c0fb0cb --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/adapter/MemoryGraphAdapterTest.java @@ -0,0 +1,744 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.adapter; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashSet; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Set; +import org.apache.geaflow.ai.graph.io.Edge; +import org.apache.geaflow.ai.graph.io.EdgeGroup; +import org.apache.geaflow.ai.graph.io.EdgeSchema; +import org.apache.geaflow.ai.graph.io.EntityGroup; +import org.apache.geaflow.ai.graph.io.MemoryGraph; +import org.apache.geaflow.ai.graph.io.Vertex; +import org.apache.geaflow.ai.graph.io.VertexGroup; +import org.apache.geaflow.ai.graph.io.VertexSchema; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.semantics.CanonicalSnapshot; +import org.apache.geaflow.ai.temporal.semantics.EventNormalizer; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +/** + * Defines the canonical MemoryGraph projection and reverse-read contract. + */ +public class MemoryGraphAdapterTest { + + private static final String CITY_VERSION = "city-version"; + private static final String CITY_TOMBSTONE = "city-tombstone"; + private static final String BOB_VERSION = "bob-version"; + private static final String BOB_COPY_VERSION = "bob-copy-version"; + private static final String CAROL_VERSION = "carol-version"; + + private static final String CITY_ADD_EVENT = "city-add"; + private static final String CITY_RETRACT_EVENT = "city-retract"; + private static final String BOB_ADD_EVENT = "bob-add"; + private static final String BOB_CORRECT_EVENT = "bob-correct"; + private static final String CAROL_ADD_EVENT = "carol-add"; + + private static final String CITY_EVIDENCE = "city-proof"; + + private static final FactKey CITY_KEY = new FactKey( + "person:alice", "city", "profile"); + private static final FactKey KNOWS_KEY = new FactKey( + "person:alice", "knows", "social"); + private static final FactKey ALIAS_KEY = new FactKey( + "person:history", "alias", "profile"); + + private final EventNormalizer normalizer = new EventNormalizer(); + + @Test + public void testEmptySnapshotCreatesFixedSchemaAndRoundTrips() { + CanonicalSnapshot empty = new CanonicalSnapshot( + new TemporalState( + Collections.emptyMap(), + Collections.emptyList()), + Collections.emptyList(), + Collections.emptyMap()); + + MemoryGraph graph = new MemoryGraphAdapter().toGraph(empty); + + Assertions.assertEquals(expectedVertexSchemas(), vertexSchemas(graph)); + Assertions.assertEquals(expectedEdgeSchemas(), edgeSchemas(graph)); + Assertions.assertEquals(expectedGroupOrder(), + new ArrayList<>(graph.entities.keySet())); + for (EntityGroup group : graph.entities.values()) { + if (group instanceof VertexGroup) { + Assertions.assertTrue( + ((VertexGroup) group).getVertices().isEmpty()); + } else { + Assertions.assertTrue( + ((EdgeGroup) group).getOutEdges().isEmpty()); + } + } + Assertions.assertEquals( + empty, + new MemoryGraphAdapter().fromGraph(graph)); + } + + @Test + public void testCompleteSnapshotRoundTripsWithoutLoss() { + CanonicalSnapshot expected = completeSnapshot(); + MemoryGraph graph = new MemoryGraphAdapter().toGraph(expected); + + CanonicalSnapshot actual = + new MemoryGraphAdapter().fromGraph(graph); + + Assertions.assertEquals(expected, actual); + Assertions.assertNotNull(graph.getVertex( + "entity", "entity:person:alice")); + Assertions.assertNotNull(graph.getVertex( + "entity", "entity:person:bob")); + Assertions.assertNotNull(graph.getVertex( + "entity", "entity:person:carol")); + Assertions.assertNotNull(graph.getVertex( + "entity", "entity:person:history")); + Assertions.assertNotNull(graph.getVertex( + "fact_version", versionId(CITY_TOMBSTONE))); + Assertions.assertNotNull(graph.getVertex( + "memory_event", eventId(CITY_RETRACT_EVENT))); + Assertions.assertNotNull(graph.getVertex( + "evidence", evidenceId(CITY_EVIDENCE))); + Assertions.assertNotNull(graph.getVertex( + "source", "source:registry")); + Assertions.assertEquals( + event(expected, CITY_ADD_EVENT).getPayloadHash(), + vertexProperty( + graph, + "memory_event", + eventId(CITY_ADD_EVENT), + "payloadHash")); + + Assertions.assertTrue(outEdges( + graph, "object", versionId(CITY_VERSION)).isEmpty()); + assertSingleEdge( + graph, + "subject", + versionId(CITY_VERSION), + "entity:person:alice"); + assertSingleEdge( + graph, + "object", + versionId(BOB_VERSION), + "entity:person:bob"); + assertSingleEdge( + graph, + "generates", + eventId(CITY_ADD_EVENT), + versionId(CITY_VERSION)); + assertSingleEdge( + graph, + "from_source", + evidenceId(CITY_EVIDENCE), + "source:registry"); + + Assertions.assertEquals( + Collections.singletonList("2"), + assertSingleEdge( + graph, + "supported_by", + eventId(CITY_ADD_EVENT), + evidenceId(CITY_EVIDENCE)).getValues()); + Assertions.assertEquals( + Collections.singletonList("2"), + assertSingleEdge( + graph, + "supported_by", + versionId(CITY_VERSION), + evidenceId(CITY_EVIDENCE)).getValues()); + + assertRelationEdge( + graph, + "supersedes", + versionId(CITY_TOMBSTONE), + versionId(CITY_VERSION), + 1); + assertRelationEdge( + graph, + "duplicate_of", + versionId(BOB_COPY_VERSION), + versionId(BOB_VERSION), + 2); + assertRelationEdge( + graph, + "conflicts_with", + versionId(BOB_VERSION), + versionId(CAROL_VERSION), + 1); + } + + @Test + public void testProjectionIsDeterministicAndUsesCanonicalEdges() { + CanonicalSnapshot snapshot = completeSnapshot(); + MemoryGraph first = new MemoryGraphAdapter().toGraph(snapshot); + MemoryGraph second = new MemoryGraphAdapter().toGraph(snapshot); + + Assertions.assertEquals(graphRows(first), graphRows(second)); + Assertions.assertEquals( + graphRows(first), + graphRows(new MemoryGraphAdapter().toGraph( + new MemoryGraphAdapter().fromGraph(first)))); + + Set identities = new HashSet<>(); + Set relationIds = new HashSet<>(); + for (VertexSchema schema + : first.getGraphSchema().getVertexSchemaList()) { + for (Vertex vertex : vertices(first, schema.getLabel())) { + Assertions.assertEquals( + schema.getFields().size(), + vertex.getValues().size()); + } + } + for (EdgeSchema schema + : first.getGraphSchema().getEdgeSchemaList()) { + for (Edge edge : edges(first, schema.getLabel())) { + Assertions.assertEquals( + schema.getFields().size(), + edge.getValues().size()); + Assertions.assertTrue( + identities.add(edge), + "duplicate edge identity: " + edge); + if (isVersionRelation(edge.getLabel())) { + Assertions.assertTrue( + relationIds.add(edge.getValues().get(0))); + } + } + } + Assertions.assertEquals(3, relationIds.size()); + assertTypedVertexIds(first); + } + + @Test + public void testNullInputsAreRejected() { + MemoryGraphAdapter adapter = new MemoryGraphAdapter(); + + Assertions.assertThrows( + NullPointerException.class, + () -> adapter.toGraph(null)); + Assertions.assertThrows( + NullPointerException.class, + () -> adapter.fromGraph(null)); + } + + @Test + public void testOrphanVerticesAreRejected() { + MemoryGraphAdapter adapter = new MemoryGraphAdapter(); + + MemoryGraph entityGraph = emptyGraph(adapter); + entityGraph.addVertex(new Vertex( + "entity", + "entity:orphan", + Collections.singletonList("person"))); + IllegalArgumentException entityError = Assertions.assertThrows( + IllegalArgumentException.class, + () -> adapter.fromGraph(entityGraph)); + Assertions.assertEquals( + "Orphan entity vertex: entity:orphan", + entityError.getMessage()); + + MemoryGraph evidenceGraph = emptyGraph(adapter); + evidenceGraph.addVertex(new Vertex( + "source", + "source:orphan-evidence-source", + Collections.singletonList("Orphan evidence source"))); + evidenceGraph.addVertex(new Vertex( + "evidence", + "evidence:orphan", + Collections.singletonList("Unreferenced evidence"))); + evidenceGraph.addEdge(new Edge( + "from_source", + "evidence:orphan", + "source:orphan-evidence-source", + Collections.emptyList())); + IllegalArgumentException evidenceError = Assertions.assertThrows( + IllegalArgumentException.class, + () -> adapter.fromGraph(evidenceGraph)); + Assertions.assertEquals( + "Orphan evidence vertex: evidence:orphan", + evidenceError.getMessage()); + + MemoryGraph sourceGraph = emptyGraph(adapter); + sourceGraph.addVertex(new Vertex( + "source", + "source:orphan", + Collections.singletonList("Unreferenced source"))); + IllegalArgumentException sourceError = Assertions.assertThrows( + IllegalArgumentException.class, + () -> adapter.fromGraph(sourceGraph)); + Assertions.assertEquals( + "Orphan source vertex: source:orphan", + sourceError.getMessage()); + } + + private static MemoryGraph emptyGraph(MemoryGraphAdapter adapter) { + return adapter.toGraph(new CanonicalSnapshot( + new TemporalState( + Collections.emptyMap(), + Collections.emptyList()), + Collections.emptyList(), + Collections.emptyMap())); + } + + private CanonicalSnapshot completeSnapshot() { + Source registry = new Source("registry", "Registry"); + Source analyst = new Source("analyst", "Analyst notes"); + Evidence cityEvidence = new Evidence( + CITY_EVIDENCE, registry, "Alice lived in Beijing"); + Evidence retractEvidence = new Evidence( + "retract-proof", registry, "The city record expired"); + Evidence bobEvidence = new Evidence( + "bob-proof", analyst, "Alice knows Bob"); + Evidence bobCopyEvidence = new Evidence( + "bob-copy-proof", analyst, "Duplicate Bob assertion"); + Evidence carolEvidence = new Evidence( + "carol-proof", registry, "Alice also knows Carol"); + Evidence aliasEvidence = new Evidence( + "alias-proof", analyst, "Historical alias event"); + + MemoryEntity alice = new MemoryEntity("person:alice", "person"); + MemoryEntity bob = new MemoryEntity("person:bob", "person"); + MemoryEntity carol = new MemoryEntity("person:carol", "person"); + MemoryEntity history = new MemoryEntity( + "person:history", "historical_person"); + MemoryFact cityFact = MemoryFact.attribute( + "fact-city", alice, "city", "Beijing"); + MemoryFact bobFact = MemoryFact.relationship( + "fact-bob", alice, "knows", bob); + MemoryFact bobCopyFact = MemoryFact.relationship( + "fact-bob-copy", alice, "knows", bob); + MemoryFact carolFact = MemoryFact.relationship( + "fact-carol", alice, "knows", carol); + MemoryFact aliasFact = MemoryFact.attribute( + "fact-alias", history, "alias", "Al"); + + TimeInterval cityValid = interval( + "2020-01-01T00:00:00Z", + "2021-01-01T00:00:00Z"); + TimeInterval relationshipValid = TimeInterval.unboundedFrom( + Instant.parse("2019-06-01T00:00:00Z")); + NormalizedMemoryEvent cityAdd = add( + CITY_ADD_EVENT, + cityFact, + CITY_KEY, + cityValid, + "2025-01-01T00:00:00Z", + cityEvidence, + cityEvidence); + NormalizedMemoryEvent cityRetract = retract( + CITY_RETRACT_EVENT, + cityFact.getId(), + CITY_KEY, + cityValid, + "2025-01-02T00:00:00Z", + retractEvidence); + NormalizedMemoryEvent bobAdd = add( + BOB_ADD_EVENT, + bobFact, + KNOWS_KEY, + relationshipValid, + "2025-01-03T00:00:00Z", + bobEvidence); + NormalizedMemoryEvent bobCorrect = correct( + BOB_CORRECT_EVENT, + bobCopyFact, + KNOWS_KEY, + relationshipValid, + "2025-01-04T00:00:00Z", + bobCopyEvidence); + NormalizedMemoryEvent carolAdd = add( + CAROL_ADD_EVENT, + carolFact, + KNOWS_KEY, + relationshipValid, + "2025-01-05T00:00:00Z", + carolEvidence); + NormalizedMemoryEvent historicalAlias = add( + "alias-history", + aliasFact, + ALIAS_KEY, + TimeInterval.unboundedFrom( + Instant.parse("2018-01-01T00:00:00Z")), + "2025-01-06T00:00:00Z", + aliasEvidence); + + MemoryFactVersion cityVersion = version( + CITY_VERSION, + cityFact, + MemoryFactVersionStatus.ACTIVE, + interval( + "2020-01-01T00:00:00.123456789Z", + "2021-01-01T00:00:00.987654321Z"), + interval( + "2025-01-01T00:00:00.000000001Z", + "2025-01-02T00:00:00.000000002Z"), + cityEvidence, + cityEvidence); + MemoryFactVersion cityTombstone = version( + CITY_TOMBSTONE, + cityFact, + MemoryFactVersionStatus.RETRACTED, + cityValid, + TimeInterval.unboundedFrom( + Instant.parse("2025-01-02T00:00:00Z")), + retractEvidence); + MemoryFactVersion bobVersion = version( + BOB_VERSION, + bobFact, + MemoryFactVersionStatus.ACTIVE, + relationshipValid, + interval( + "2025-01-03T00:00:00Z", + "2025-01-04T00:00:00Z"), + bobEvidence); + MemoryFactVersion bobCopyVersion = version( + BOB_COPY_VERSION, + bobCopyFact, + MemoryFactVersionStatus.ACTIVE, + relationshipValid, + interval( + "2025-01-04T00:00:00Z", + "2025-01-05T00:00:00Z"), + bobCopyEvidence); + MemoryFactVersion carolVersion = version( + CAROL_VERSION, + carolFact, + MemoryFactVersionStatus.ACTIVE, + relationshipValid, + TimeInterval.unboundedFrom( + Instant.parse("2025-01-05T00:00:00Z")), + carolEvidence); + + Map> versions = + new LinkedHashMap<>(); + versions.put(KNOWS_KEY, Arrays.asList( + carolVersion, bobCopyVersion, bobVersion)); + versions.put(CITY_KEY, Arrays.asList( + cityTombstone, cityVersion)); + + VersionRelation duplicate = new VersionRelation( + VersionRelationType.DUPLICATE_OF, + BOB_COPY_VERSION, + BOB_VERSION); + List relations = Arrays.asList( + new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + BOB_VERSION, + CAROL_VERSION), + duplicate, + new VersionRelation( + VersionRelationType.SUPERSEDES, + CITY_TOMBSTONE, + CITY_VERSION), + duplicate); + + Map generating = new LinkedHashMap<>(); + generating.put(CAROL_VERSION, CAROL_ADD_EVENT); + generating.put(CITY_TOMBSTONE, CITY_RETRACT_EVENT); + generating.put(BOB_COPY_VERSION, BOB_CORRECT_EVENT); + generating.put(CITY_VERSION, CITY_ADD_EVENT); + generating.put(BOB_VERSION, BOB_ADD_EVENT); + + return new CanonicalSnapshot( + new TemporalState(versions, relations), + Arrays.asList( + historicalAlias, + carolAdd, + bobCorrect, + bobAdd, + cityRetract, + cityAdd), + generating); + } + + private NormalizedMemoryEvent add( + String eventId, + MemoryFact fact, + FactKey key, + TimeInterval validTime, + String recordedAt, + Evidence... evidence) { + return normalizer.normalize( + MemoryEvent.add( + eventId, + fact, + validTime, + Instant.parse(recordedAt), + Arrays.asList(evidence)), + key); + } + + private NormalizedMemoryEvent correct( + String eventId, + MemoryFact fact, + FactKey key, + TimeInterval validTime, + String recordedAt, + Evidence... evidence) { + return normalizer.normalize( + MemoryEvent.correct( + eventId, + fact, + validTime, + Instant.parse(recordedAt), + Arrays.asList(evidence)), + key); + } + + private NormalizedMemoryEvent retract( + String eventId, + String factId, + FactKey key, + TimeInterval validTime, + String recordedAt, + Evidence... evidence) { + return normalizer.normalize( + MemoryEvent.retract( + eventId, + factId, + validTime, + Instant.parse(recordedAt), + Arrays.asList(evidence)), + key); + } + + private static MemoryFactVersion version( + String versionId, + MemoryFact fact, + MemoryFactVersionStatus status, + TimeInterval validTime, + TimeInterval transactionTime, + Evidence... evidence) { + return new MemoryFactVersion( + versionId, + fact, + status, + validTime, + transactionTime, + Arrays.asList(evidence)); + } + + private static TimeInterval interval(String start, String end) { + return new TimeInterval(Instant.parse(start), Instant.parse(end)); + } + + private static List expectedVertexSchemas() { + return Arrays.asList( + "entity|id|[label]", + "fact_version|id|[factId, predicate, scope, valueKind, " + + "literalValue, status, validStart, validEnd, " + + "transactionStart, transactionEnd]", + "memory_event|id|[operation, factId, subjectId, predicate, " + + "scope, valueKind, value, validStart, validEnd, " + + "recordedAt, payloadHash]", + "evidence|id|[content]", + "source|id|[name]"); + } + + private static List expectedEdgeSchemas() { + return Arrays.asList( + "subject|srcId|dstId|[]", + "object|srcId|dstId|[]", + "generates|srcId|dstId|[]", + "supported_by|srcId|dstId|[occurrenceCount]", + "from_source|srcId|dstId|[]", + "supersedes|srcId|dstId|[relationId, occurrenceCount]", + "duplicate_of|srcId|dstId|[relationId, occurrenceCount]", + "conflicts_with|srcId|dstId|[relationId, occurrenceCount]"); + } + + private static List expectedGroupOrder() { + return Arrays.asList( + "entity", + "fact_version", + "memory_event", + "evidence", + "source", + "subject", + "object", + "generates", + "supported_by", + "from_source", + "supersedes", + "duplicate_of", + "conflicts_with"); + } + + private static List vertexSchemas(MemoryGraph graph) { + List rows = new ArrayList<>(); + for (VertexSchema schema + : graph.getGraphSchema().getVertexSchemaList()) { + rows.add(schema.getLabel() + "|" + schema.getIdField() + + "|" + schema.getFields()); + } + return rows; + } + + private static List edgeSchemas(MemoryGraph graph) { + List rows = new ArrayList<>(); + for (EdgeSchema schema + : graph.getGraphSchema().getEdgeSchemaList()) { + rows.add(schema.getLabel() + "|" + schema.getSrcIdField() + + "|" + schema.getDstIdField() + + "|" + schema.getFields()); + } + return rows; + } + + private static List graphRows(MemoryGraph graph) { + List rows = new ArrayList<>(); + rows.addAll(vertexSchemas(graph)); + rows.addAll(edgeSchemas(graph)); + for (VertexSchema schema + : graph.getGraphSchema().getVertexSchemaList()) { + for (Vertex vertex : vertices(graph, schema.getLabel())) { + rows.add("vertex|" + vertex.getLabel() + "|" + + vertex.getId() + "|" + vertex.getValues()); + } + } + for (EdgeSchema schema + : graph.getGraphSchema().getEdgeSchemaList()) { + for (Edge edge : edges(graph, schema.getLabel())) { + rows.add("edge|" + edge.getLabel() + "|" + + edge.getSrcId() + "|" + edge.getDstId() + + "|" + edge.getValues()); + } + } + return rows; + } + + private static List vertices(MemoryGraph graph, String label) { + return ((VertexGroup) graph.entities.get(label)).getVertices(); + } + + private static String vertexProperty( + MemoryGraph graph, + String label, + String vertexId, + String field) { + Vertex vertex = graph.getVertex(label, vertexId); + Assertions.assertNotNull(vertex); + for (VertexSchema schema + : graph.getGraphSchema().getVertexSchemaList()) { + if (label.equals(schema.getLabel())) { + int index = schema.getFields().indexOf(field); + Assertions.assertTrue(index >= 0); + return vertex.getValues().get(index); + } + } + throw new AssertionError("Missing vertex schema: " + label); + } + + private static List edges(MemoryGraph graph, String label) { + return ((EdgeGroup) graph.entities.get(label)).getOutEdges(); + } + + private static List outEdges( + MemoryGraph graph, + String label, + String sourceId) { + return ((EdgeGroup) graph.entities.get(label)).getOutEdges(sourceId); + } + + private static Edge assertSingleEdge( + MemoryGraph graph, + String label, + String sourceId, + String targetId) { + List matches = graph.getEdge(label, sourceId, targetId); + Assertions.assertEquals(1, matches.size()); + return matches.get(0); + } + + private static void assertRelationEdge( + MemoryGraph graph, + String label, + String sourceId, + String targetId, + int occurrenceCount) { + List values = assertSingleEdge( + graph, label, sourceId, targetId).getValues(); + Assertions.assertEquals(2, values.size()); + Assertions.assertFalse(values.get(0).trim().isEmpty()); + Assertions.assertEquals( + Integer.toString(occurrenceCount), values.get(1)); + } + + private static boolean isVersionRelation(String label) { + return "supersedes".equals(label) + || "duplicate_of".equals(label) + || "conflicts_with".equals(label); + } + + private static void assertTypedVertexIds(MemoryGraph graph) { + Map prefixes = new LinkedHashMap<>(); + prefixes.put("entity", "entity:"); + prefixes.put("fact_version", "version:"); + prefixes.put("memory_event", "event:"); + prefixes.put("evidence", "evidence:"); + prefixes.put("source", "source:"); + for (Map.Entry entry : prefixes.entrySet()) { + for (Vertex vertex : vertices(graph, entry.getKey())) { + Assertions.assertTrue( + vertex.getId().startsWith(entry.getValue())); + } + } + } + + private static String versionId(String rawId) { + return "version:" + rawId; + } + + private static String eventId(String rawId) { + return "event:" + rawId; + } + + private static String evidenceId(String rawId) { + return "evidence:" + rawId; + } + + private static NormalizedMemoryEvent event( + CanonicalSnapshot snapshot, + String eventId) { + for (NormalizedMemoryEvent event : snapshot.getEvents()) { + if (eventId.equals(event.getEventId())) { + return event; + } + } + throw new AssertionError("Missing event: " + eventId); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/baseline/BaselineReplayTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/baseline/BaselineReplayTest.java new file mode 100644 index 000000000..f3eefc5d3 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/baseline/BaselineReplayTest.java @@ -0,0 +1,519 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.baseline; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.oracle.ReplayMethod; +import org.apache.geaflow.ai.temporal.semantics.CanonicalSnapshot; +import org.apache.geaflow.ai.temporal.semantics.EventNormalizer; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +/** + * Defines deterministic current-state semantics for the two replay baselines. + */ +public class BaselineReplayTest { + + private static final FactKey PROFILE_KEY = new FactKey( + "person:alice", + "city", + "profile"); + private static final FactKey ACCOUNT_KEY = new FactKey( + "person:alice", + "city", + "account"); + private static final TimeInterval ALL_VALID_TIME = + TimeInterval.unboundedFrom(Instant.ofEpochMilli(Long.MIN_VALUE)); + + private final EventNormalizer normalizer = new EventNormalizer(); + + @Test + public void testSameTimeConflictUsesEventIdTieBreakAndIgnoresValidTime() { + Instant recordedAt = Instant.ofEpochMilli(1000); + NormalizedMemoryEvent eventA = add( + "event-a", + "fact-a", + PROFILE_KEY, + "Beijing", + interval(100, 200), + recordedAt); + NormalizedMemoryEvent eventB = add( + "event-b", + "fact-b", + PROFILE_KEY, + "Rome", + interval(300, 400), + recordedAt); + List forward = Arrays.asList(eventA, eventB); + List reversed = Arrays.asList(eventB, eventA); + + ReplayMethod lww = new LwwBaseline(); + CanonicalSnapshot expectedLww = snapshot( + forward, + versions(PROFILE_KEY, version(eventB)), + eventB); + Assertions.assertEquals(expectedLww, lww.replayToSnapshot(forward)); + Assertions.assertEquals(expectedLww, lww.replayToSnapshot(reversed)); + + ReplayMethod singleTimestamp = new SingleTimestampBaseline(); + CanonicalSnapshot expectedSingle = snapshot( + forward, + versions(PROFILE_KEY, version(eventA), version(eventB)), + Collections.singletonList(new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + versionId(eventA), + versionId(eventB))), + eventA, + eventB); + Assertions.assertEquals( + expectedSingle, + singleTimestamp.replayToSnapshot(forward)); + Assertions.assertEquals( + expectedSingle, + singleTimestamp.replayToSnapshot(reversed)); + } + + @Test + public void testSingleTimestampReplacesSameValueAndBuildsAllCurrentConflicts() { + NormalizedMemoryEvent firstBeijing = add( + "event-a", + "fact-a", + PROFILE_KEY, + "Beijing", + interval(100, 200), + Instant.ofEpochMilli(1000)); + NormalizedMemoryEvent latestBeijing = add( + "event-b", + "fact-b", + PROFILE_KEY, + "Beijing", + interval(300, 400), + Instant.ofEpochMilli(2000)); + NormalizedMemoryEvent rome = add( + "event-c", + "fact-c", + PROFILE_KEY, + "Rome", + interval(500, 600), + Instant.ofEpochMilli(3000)); + NormalizedMemoryEvent paris = add( + "event-d", + "fact-d", + PROFILE_KEY, + "Paris", + interval(700, 800), + Instant.ofEpochMilli(4000)); + List ordered = Arrays.asList( + firstBeijing, + latestBeijing, + rome, + paris); + List reversed = Arrays.asList( + paris, + rome, + latestBeijing, + firstBeijing); + + ReplayMethod method = new SingleTimestampBaseline(); + CanonicalSnapshot snapshot = method.replayToSnapshot(ordered); + + Assertions.assertEquals(snapshot, method.replayToSnapshot(reversed)); + Assertions.assertEquals( + Arrays.asList( + versionId(latestBeijing), + versionId(rome), + versionId(paris)), + versionIds(snapshot)); + Assertions.assertEquals( + latestBeijing.getEventId(), + snapshot.getGeneratingEventIds().get( + versionId(latestBeijing))); + Assertions.assertEquals( + Arrays.asList( + conflict(latestBeijing, rome), + conflict(latestBeijing, paris), + conflict(rome, paris)), + snapshot.getState().getRelations()); + } + + @Test + public void testLatePartialCorrectionReplacesCurrentStateWithoutHistory() { + NormalizedMemoryEvent first = add( + "event-add-a", + "fact-a", + PROFILE_KEY, + "Beijing", + interval(100, 900), + Instant.ofEpochMilli(1000)); + NormalizedMemoryEvent conflict = add( + "event-add-b", + "fact-b", + PROFILE_KEY, + "Rome", + interval(100, 900), + Instant.ofEpochMilli(2000)); + NormalizedMemoryEvent correction = correct( + "event-correct", + "fact-a", + PROFILE_KEY, + "Shanghai", + interval(400, 500), + Instant.ofEpochMilli(3000)); + List chronological = + Arrays.asList(first, conflict, correction); + List arrivalOrder = + Arrays.asList(correction, conflict, first); + CanonicalSnapshot expected = snapshot( + chronological, + versions(PROFILE_KEY, version(correction)), + correction); + + for (ReplayMethod method : methods()) { + Assertions.assertEquals( + expected, + method.replayToSnapshot(chronological)); + Assertions.assertEquals( + expected, + method.replayToSnapshot(arrivalOrder)); + } + } + + @Test + public void testRetractClearsCurrentStateAndLaterAddRestoresIt() { + NormalizedMemoryEvent first = add( + "event-add-a", + "fact-a", + PROFILE_KEY, + "Beijing", + interval(100, 900), + Instant.ofEpochMilli(1000)); + NormalizedMemoryEvent second = add( + "event-add-b", + "fact-b", + PROFILE_KEY, + "Rome", + interval(100, 900), + Instant.ofEpochMilli(2000)); + NormalizedMemoryEvent retract = retract( + "event-retract", + "fact-b", + PROFILE_KEY, + interval(450, 460), + Instant.ofEpochMilli(3000)); + NormalizedMemoryEvent recovery = add( + "event-recover", + "fact-c", + PROFILE_KEY, + "Paris", + interval(700, 800), + Instant.ofEpochMilli(4000)); + List throughRetract = + Arrays.asList(retract, second, first); + CanonicalSnapshot expectedRetracted = snapshot( + throughRetract, + Collections.emptyMap()); + List throughRecovery = + Arrays.asList(recovery, retract, second, first); + CanonicalSnapshot expectedRecovered = snapshot( + throughRecovery, + versions(PROFILE_KEY, version(recovery)), + recovery); + NormalizedMemoryEvent correctionWithoutCurrent = correct( + "event-orphan-correct", + "fact-a", + PROFILE_KEY, + "London", + interval(100, 200), + Instant.ofEpochMilli(1000)); + NormalizedMemoryEvent retractWithoutCurrent = retract( + "event-orphan-retract", + "fact-a", + PROFILE_KEY, + interval(100, 200), + Instant.ofEpochMilli(1000)); + + for (ReplayMethod method : methods()) { + Assertions.assertEquals( + expectedRetracted, + method.replayToSnapshot(throughRetract)); + Assertions.assertEquals( + expectedRecovered, + method.replayToSnapshot(throughRecovery)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> method.replayToSnapshot( + Collections.singletonList(correctionWithoutCurrent))); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> method.replayToSnapshot( + Collections.singletonList(retractWithoutCurrent))); + Assertions.assertEquals( + expectedRetracted, + method.replayToSnapshot(throughRetract)); + } + } + + @Test + public void testDuplicateIsNoopAndEventIdReuseIsAtomic() { + NormalizedMemoryEvent original = add( + "event-a", + "fact-a", + PROFILE_KEY, + "Beijing", + interval(100, 900), + Instant.ofEpochMilli(1000)); + NormalizedMemoryEvent reused = add( + "event-a", + "fact-a", + PROFILE_KEY, + "Rome", + interval(100, 900), + Instant.ofEpochMilli(1000)); + NormalizedMemoryEvent retract = retract( + "event-retract", + "fact-a", + PROFILE_KEY, + interval(200, 300), + Instant.ofEpochMilli(2000)); + CanonicalSnapshot expected = snapshot( + Collections.singletonList(original), + versions(PROFILE_KEY, version(original)), + original); + CanonicalSnapshot expectedRetracted = snapshot( + Arrays.asList(original, retract), + Collections.emptyMap()); + + for (ReplayMethod method : methods()) { + Assertions.assertEquals( + expected, + method.replayToSnapshot( + Arrays.asList(original, original))); + Assertions.assertEquals( + expectedRetracted, + method.replayToSnapshot( + Arrays.asList(original, retract, retract))); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> method.replayToSnapshot( + Arrays.asList(original, reused))); + Assertions.assertEquals( + expected, + method.replayToSnapshot( + Collections.singletonList(original))); + } + } + + @Test + public void testFactKeyScopeKeepsIndependentCurrentState() { + Instant recordedAt = Instant.ofEpochMilli(1000); + NormalizedMemoryEvent profile = add( + "event-profile", + "fact-shared", + PROFILE_KEY, + "Beijing", + interval(100, 200), + recordedAt); + NormalizedMemoryEvent account = add( + "event-account", + "fact-shared", + ACCOUNT_KEY, + "Rome", + interval(300, 400), + recordedAt); + Map> current = new LinkedHashMap<>(); + current.put(PROFILE_KEY, Collections.singletonList(version(profile))); + current.put(ACCOUNT_KEY, Collections.singletonList(version(account))); + CanonicalSnapshot expected = snapshot( + Arrays.asList(profile, account), + current, + profile, + account); + + for (ReplayMethod method : methods()) { + Assertions.assertEquals( + expected, + method.replayToSnapshot(Arrays.asList(profile, account))); + } + } + + private static ReplayMethod[] methods() { + return new ReplayMethod[] { + new LwwBaseline(), + new SingleTimestampBaseline() + }; + } + + private NormalizedMemoryEvent add( + String eventId, + String factId, + FactKey key, + String value, + TimeInterval validTime, + Instant recordedAt) { + return normalizer.normalize( + MemoryEvent.add( + eventId, + fact(factId, key, value), + validTime, + recordedAt, + evidence(eventId)), + key); + } + + private NormalizedMemoryEvent correct( + String eventId, + String factId, + FactKey key, + String value, + TimeInterval validTime, + Instant recordedAt) { + return normalizer.normalize( + MemoryEvent.correct( + eventId, + fact(factId, key, value), + validTime, + recordedAt, + evidence(eventId)), + key); + } + + private NormalizedMemoryEvent retract( + String eventId, + String factId, + FactKey key, + TimeInterval validTime, + Instant recordedAt) { + return normalizer.normalize( + MemoryEvent.retract( + eventId, + factId, + validTime, + recordedAt, + evidence(eventId)), + key); + } + + private static MemoryFact fact( + String factId, + FactKey key, + String value) { + return MemoryFact.attribute( + factId, + new MemoryEntity(key.getSubjectId(), "person"), + key.getPredicate(), + value); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-registry", "registry"), + "Evidence for " + eventId)); + } + + private static TimeInterval interval(long start, long end) { + return new TimeInterval( + Instant.ofEpochMilli(start), + Instant.ofEpochMilli(end)); + } + + private static MemoryFactVersion version( + NormalizedMemoryEvent event) { + return new MemoryFactVersion( + versionId(event), + event.getEvent().getFact().get(), + MemoryFactVersionStatus.ACTIVE, + ALL_VALID_TIME, + TimeInterval.unboundedFrom(event.getRecordedAt()), + event.getEvidence()); + } + + private static String versionId(NormalizedMemoryEvent event) { + return event.getEventId() + ":version:0"; + } + + private static List versionIds(CanonicalSnapshot snapshot) { + List ids = new ArrayList<>(); + for (MemoryFactVersion version : snapshot.getState().getVersions()) { + ids.add(version.getId()); + } + return ids; + } + + private static VersionRelation conflict( + NormalizedMemoryEvent left, + NormalizedMemoryEvent right) { + return new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + versionId(left), + versionId(right)); + } + + private static Map> versions( + FactKey key, + MemoryFactVersion... values) { + return Collections.singletonMap(key, Arrays.asList(values)); + } + + private static CanonicalSnapshot snapshot( + List events, + Map> versions, + NormalizedMemoryEvent... generatingEvents) { + return snapshot( + events, + versions, + Collections.emptyList(), + generatingEvents); + } + + private static CanonicalSnapshot snapshot( + List events, + Map> versions, + List relations, + NormalizedMemoryEvent... generatingEvents) { + Map generating = new LinkedHashMap<>(); + for (NormalizedMemoryEvent event : generatingEvents) { + generating.put(versionId(event), event.getEventId()); + } + return new CanonicalSnapshot( + new TemporalState(versions, relations), + events, + generating); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalIntegratorTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalIntegratorTest.java new file mode 100644 index 000000000..9c1c87ce1 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalIntegratorTest.java @@ -0,0 +1,385 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.oracle.FullReplayOracle; +import org.apache.geaflow.ai.temporal.query.BitemporalQuery; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class IncrementalTemporalIntegratorTest { + + private final IncrementalTemporalIntegrator integrator = + new IncrementalTemporalIntegrator(); + private final FullReplayOracle oracle = new FullReplayOracle(); + private final BitemporalQuery query = new BitemporalQuery(); + + @Test + public void testOrderedEventsMatchFullReplayAfterEachApply() { + List events = Arrays.asList( + addEvent( + "event-add", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"), + correctEvent( + "event-correct", + "fact-alice-city", + "person:alice", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + retractEvent( + "event-retract", + "fact-alice-city", + "2024-08-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-11-01T00:00:00Z")); + + assertMatchesAfterEachApply(events); + } + + @Test + public void testLateEventMatchesFullReplayAndKeepsOtherFact() { + MemoryEvent aliceAdd = addEvent( + "event-alice-add", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent bobAdd = addEvent( + "event-bob-add", + "fact-bob-city", + "person:bob", + "Paris", + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"); + MemoryEvent correction = correctEvent( + "event-alice-correct", + "fact-alice-city", + "person:alice", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + MemoryEvent retract = retractEvent( + "event-alice-retract", + "fact-alice-city", + "2024-08-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-11-01T00:00:00Z"); + MemoryEvent lateCorrection = correctEvent( + "event-alice-late", + "fact-alice-city", + "person:alice", + "Tianjin", + "2024-02-01T00:00:00Z", + "2024-03-01T00:00:00Z", + "2024-05-15T00:00:00Z"); + + List arrivalOrder = Arrays.asList( + aliceAdd, + bobAdd, + correction, + retract, + lateCorrection); + assertMatchesAfterEachApply(arrivalOrder); + + List visible = query.query( + integrator.snapshot(), + time("2024-02-15T00:00:00Z"), + time("2024-12-01T00:00:00Z")); + + Assertions.assertEquals( + lateCorrection.getFact().get(), + findVersion(visible, "fact-alice-city").getFact()); + Assertions.assertEquals( + bobAdd.getFact().get(), + findVersion(visible, "fact-bob-city").getFact()); + } + + @Test + public void testDuplicateAndConflictingEventIdAreAtomic() { + MemoryEvent event = addEvent( + "event-add", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent duplicate = addEvent( + "event-add", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent conflict = addEvent( + "event-add", + "fact-alice-city", + "person:alice", + "Shanghai", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + + integrator.apply(event); + List expected = integrator.snapshot(); + + integrator.apply(duplicate); + Assertions.assertEquals(expected, integrator.snapshot()); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> integrator.apply(conflict)); + Assertions.assertEquals(expected, integrator.snapshot()); + } + + @Test + public void testInvalidEventDoesNotChangeState() { + MemoryEvent add = addEvent( + "event-add", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent invalid = correctEvent( + "event-change", + "fact-alice-city", + "person:alice", + "Shanghai", + "2023-12-01T00:00:00Z", + "2024-02-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + MemoryEvent retry = correctEvent( + "event-change", + "fact-alice-city", + "person:alice", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + + integrator.apply(add); + List before = integrator.snapshot(); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> integrator.apply(invalid)); + Assertions.assertEquals(before, integrator.snapshot()); + + integrator.apply(retry); + Assertions.assertEquals( + oracle.replay(Arrays.asList(add, retry)), + integrator.snapshot()); + } + + @Test + public void testSnapshotIsDeterministicAndImmutable() { + MemoryEvent bob = addEvent( + "event-b", + "fact-bob-city", + "person:bob", + "Paris", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent alice = addEvent( + "event-a", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + + integrator.apply(bob); + integrator.apply(alice); + List snapshot = integrator.snapshot(); + + Assertions.assertEquals(2, snapshot.size()); + Assertions.assertEquals( + "event-a:version:0", + snapshot.get(0).getId()); + Assertions.assertEquals( + "event-b:version:0", + snapshot.get(1).getId()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.clear()); + } + + @Test + public void testEventSnapshotIsDeterministicAndImmutable() { + MemoryEvent tieLater = addEvent( + "event-b", + "fact-bob-city", + "person:bob", + "Paris", + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"); + MemoryEvent early = addEvent( + "event-c", + "fact-carol-city", + "person:carol", + "Rome", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent tieEarlier = addEvent( + "event-a", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"); + + integrator.apply(tieLater); + integrator.apply(early); + integrator.apply(tieEarlier); + integrator.apply(tieLater); + List snapshot = integrator.eventSnapshot(); + + Assertions.assertEquals( + Arrays.asList(early, tieEarlier, tieLater), + snapshot); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.clear()); + } + + @Test + public void testEmptyAndNullInput() { + Assertions.assertTrue(integrator.snapshot().isEmpty()); + Assertions.assertThrows( + NullPointerException.class, + () -> integrator.apply((MemoryEvent) null)); + Assertions.assertTrue(integrator.snapshot().isEmpty()); + } + + private void assertMatchesAfterEachApply( + List arrivalOrder) { + List received = new ArrayList<>(); + for (MemoryEvent event : arrivalOrder) { + integrator.apply(event); + received.add(event); + Assertions.assertEquals( + oracle.replay(received), + integrator.snapshot()); + } + } + + private static MemoryFactVersion findVersion( + List versions, + String factId) { + for (MemoryFactVersion version : versions) { + if (version.getFact().getId().equals(factId)) { + return version; + } + } + throw new AssertionError( + "Missing fact version: " + factId); + } + + private static MemoryEvent addEvent( + String eventId, + String factId, + String subjectId, + String literalValue, + String validStart, + String transactionTime) { + return MemoryEvent.add( + eventId, + fact(factId, subjectId, literalValue), + TimeInterval.unboundedFrom(time(validStart)), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent correctEvent( + String eventId, + String factId, + String subjectId, + String literalValue, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.correct( + eventId, + fact(factId, subjectId, literalValue), + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent retractEvent( + String eventId, + String factId, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.retract( + eventId, + factId, + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryFact fact( + String factId, + String subjectId, + String literalValue) { + return MemoryFact.attribute( + factId, + new MemoryEntity(subjectId, "person"), + "city", + literalValue); + } + + private static TimeInterval interval( + String start, + String end) { + return new TimeInterval(time(start), time(end)); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "customer-database"), + "Evidence for " + eventId)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalStateTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalStateTest.java new file mode 100644 index 000000000..1122e5c7b --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/IncrementalTemporalStateTest.java @@ -0,0 +1,405 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.oracle.FullReplayOracle; +import org.apache.geaflow.ai.temporal.semantics.EventNormalizer; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class IncrementalTemporalStateTest { + + private static final FactKey ALICE_KEY = new FactKey( + "person:alice", + "city", + "profile"); + private static final FactKey BOB_KEY = new FactKey( + "person:bob", + "city", + "profile"); + + private final EventNormalizer normalizer = new EventNormalizer(); + private final FullReplayOracle oracle = new FullReplayOracle(); + + @Test + public void testLateEventMatchesOracleAndKeepsOtherFactKey() { + IncrementalTemporalIntegrator integrator = + new IncrementalTemporalIntegrator(); + TimeInterval validTime = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + NormalizedMemoryEvent aliceAdd = add( + "event-alice-add", + ALICE_KEY, + "Beijing", + validTime, + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent bobAdd = add( + "event-bob-add", + BOB_KEY, + "Paris", + validTime, + "2024-04-01T00:00:00Z"); + NormalizedMemoryEvent laterCorrection = correct( + "event-alice-later", + ALICE_KEY, + "Rome", + validTime, + "2024-08-01T00:00:00Z"); + NormalizedMemoryEvent retraction = retract( + "event-alice-retract", + ALICE_KEY, + validTime, + "2024-10-01T00:00:00Z"); + NormalizedMemoryEvent lateCorrection = correct( + "event-alice-late", + ALICE_KEY, + "Shanghai", + validTime, + "2024-06-01T00:00:00Z"); + List received = new ArrayList<>(); + + for (NormalizedMemoryEvent event : Arrays.asList( + aliceAdd, + bobAdd, + laterCorrection, + retraction)) { + integrator.apply(event); + received.add(event); + Assertions.assertEquals( + oracle.replayNormalized(received), + integrator.stateSnapshot()); + } + + List bobBefore = + integrator.stateSnapshot() + .getVersionsByFactKey().get(BOB_KEY); + integrator.apply(lateCorrection); + received.add(lateCorrection); + + TemporalState state = integrator.stateSnapshot(); + Assertions.assertEquals( + oracle.replayNormalized(received), + state); + Assertions.assertEquals( + bobBefore, + state.getVersionsByFactKey().get(BOB_KEY)); + Assertions.assertEquals( + MemoryFactVersionStatus.RETRACTED, + version(state, "event-alice-retract:version:0") + .getStatus()); + } + + @Test + public void testPartialChangesAddConflictsForResidualVersions() { + IncrementalTemporalIntegrator integrator = + new IncrementalTemporalIntegrator(); + NormalizedMemoryEvent first = add( + "event-first", + ALICE_KEY, + "Beijing", + interval( + "2024-01-01T00:00:00Z", + "2024-10-01T00:00:00Z"), + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent conflict = add( + "event-conflict", + ALICE_KEY, + "Shanghai", + interval( + "2024-05-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + "2024-04-01T00:00:00Z"); + NormalizedMemoryEvent correction = correct( + "event-correct", + ALICE_KEY, + "Rome", + interval( + "2024-01-01T00:00:00Z", + "2024-05-01T00:00:00Z"), + "2024-06-01T00:00:00Z"); + NormalizedMemoryEvent retraction = retract( + "event-retract", + ALICE_KEY, + correction.getValidTime(), + "2024-06-01T00:00:00Z"); + + integrator.apply(first); + integrator.apply(conflict); + integrator.apply(correction); + + List relations = + integrator.stateSnapshot().getRelations(); + Assertions.assertTrue(relations.contains(new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + "event-conflict:version:0", + "event-first:version:0"))); + Assertions.assertTrue(relations.contains(new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + "event-conflict:version:0", + "event-correct:version:1"))); + + IncrementalTemporalIntegrator retractingIntegrator = + new IncrementalTemporalIntegrator(); + retractingIntegrator.apply(first); + retractingIntegrator.apply(conflict); + retractingIntegrator.apply(retraction); + Assertions.assertTrue( + retractingIntegrator.stateSnapshot().getRelations().contains( + new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + "event-conflict:version:0", + "event-retract:version:1"))); + } + + @Test + public void testFailedApplyDoesNotCommitLedgerOrState() { + IncrementalTemporalIntegrator integrator = + new IncrementalTemporalIntegrator(); + TimeInterval validTime = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + NormalizedMemoryEvent add = add( + "event-add", + ALICE_KEY, + "Beijing", + validTime, + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent invalid = correct( + "event-change", + ALICE_KEY, + "Shanghai", + interval( + "2023-12-01T00:00:00Z", + "2024-02-01T00:00:00Z"), + "2024-06-01T00:00:00Z"); + NormalizedMemoryEvent validRetry = correct( + "event-change", + ALICE_KEY, + "Shanghai", + validTime, + "2024-06-01T00:00:00Z"); + NormalizedMemoryEvent reused = correct( + "event-change", + ALICE_KEY, + "London", + validTime, + "2024-06-01T00:00:00Z"); + + integrator.apply(add); + TemporalState before = integrator.stateSnapshot(); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> integrator.apply(invalid)); + Assertions.assertEquals(before, integrator.stateSnapshot()); + + integrator.apply(validRetry); + TemporalState afterRetry = integrator.stateSnapshot(); + Assertions.assertEquals( + oracle.replayNormalized(Arrays.asList(add, validRetry)), + afterRetry); + + integrator.apply(validRetry); + Assertions.assertEquals(afterRetry, integrator.stateSnapshot()); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> integrator.apply(reused)); + Assertions.assertEquals(afterRetry, integrator.stateSnapshot()); + } + + @Test + public void testFactKeyScopeIsolationAndImmutableSnapshot() { + IncrementalTemporalIntegrator integrator = + new IncrementalTemporalIntegrator(); + FactKey accountKey = new FactKey( + "person:alice", + "city", + "account"); + TimeInterval validTime = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + NormalizedMemoryEvent profile = add( + "event-profile", + ALICE_KEY, + "Beijing", + "fact-shared", + validTime, + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent account = add( + "event-account", + accountKey, + "Shanghai", + "fact-shared", + validTime, + "2024-04-01T00:00:00Z"); + + integrator.apply(profile); + integrator.apply(account); + TemporalState state = integrator.stateSnapshot(); + + Assertions.assertEquals(2, state.getVersions().size()); + Assertions.assertEquals( + Arrays.asList(accountKey, ALICE_KEY), + new ArrayList<>( + state.getVersionsByFactKey().keySet())); + Assertions.assertTrue(state.getRelations().isEmpty()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> state.getVersions().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> state.getRelations().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> state.getVersionsByFactKey().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> state.getVersionsByFactKey() + .get(ALICE_KEY).clear()); + } + + private NormalizedMemoryEvent add( + String eventId, + FactKey key, + String value, + TimeInterval validTime, + String recordedAt) { + return add( + eventId, + key, + value, + factId(key), + validTime, + recordedAt); + } + + private NormalizedMemoryEvent add( + String eventId, + FactKey key, + String value, + String factId, + TimeInterval validTime, + String recordedAt) { + return normalizer.normalize( + MemoryEvent.add( + eventId, + fact(factId, key, value), + validTime, + time(recordedAt), + evidence(eventId)), + key); + } + + private NormalizedMemoryEvent correct( + String eventId, + FactKey key, + String value, + TimeInterval validTime, + String recordedAt) { + return normalizer.normalize( + MemoryEvent.correct( + eventId, + fact(factId(key), key, value), + validTime, + time(recordedAt), + evidence(eventId)), + key); + } + + private NormalizedMemoryEvent retract( + String eventId, + FactKey key, + TimeInterval validTime, + String recordedAt) { + return normalizer.normalize( + MemoryEvent.retract( + eventId, + factId(key), + validTime, + time(recordedAt), + evidence(eventId)), + key); + } + + private static MemoryFact fact( + String factId, + FactKey key, + String value) { + return MemoryFact.attribute( + factId, + new MemoryEntity(key.getSubjectId(), "person"), + key.getPredicate(), + value); + } + + private static String factId(FactKey key) { + return "fact-" + + key.getSubjectId() + + "-" + + key.getScope(); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "registry"), + "Evidence for " + eventId)); + } + + private static MemoryFactVersion version( + TemporalState state, + String versionId) { + for (MemoryFactVersion version : state.getVersions()) { + if (version.getId().equals(versionId)) { + return version; + } + } + throw new AssertionError("Missing version: " + versionId); + } + + private static TimeInterval interval( + String start, + String end) { + return new TimeInterval(time(start), time(end)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalEventAggregateFunctionTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalEventAggregateFunctionTest.java new file mode 100644 index 000000000..97c90cfff --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalEventAggregateFunctionTest.java @@ -0,0 +1,276 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.oracle.FullReplayOracle; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class TemporalEventAggregateFunctionTest { + + private final TemporalEventAggregateFunction function = + new TemporalEventAggregateFunction(); + private final FullReplayOracle oracle = new FullReplayOracle(); + + @Test + public void testCreateAddAndGetResultMatchFullReplay() { + IncrementalTemporalIntegrator accumulator = + function.createAccumulator(); + IncrementalTemporalIntegrator other = + function.createAccumulator(); + List events = Arrays.asList( + addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z"), + correctEvent( + "event-correct", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + retractEvent( + "event-retract", + "2024-08-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-11-01T00:00:00Z"), + correctEvent( + "event-late", + "Tianjin", + "2024-02-01T00:00:00Z", + "2024-03-01T00:00:00Z", + "2024-05-01T00:00:00Z")); + + Assertions.assertNotSame(accumulator, other); + Assertions.assertTrue(function.getResult(other).isEmpty()); + + List received = new ArrayList<>(); + for (MemoryEvent event : events) { + function.add(event, accumulator); + received.add(event); + Assertions.assertEquals( + oracle.replay(received), + function.getResult(accumulator)); + } + + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> function.getResult(accumulator).clear()); + } + + @Test + public void testDuplicateAndConflictRemainAtomic() { + IncrementalTemporalIntegrator accumulator = + function.createAccumulator(); + MemoryEvent event = addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z"); + MemoryEvent duplicate = addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z"); + MemoryEvent conflict = addEvent( + "event-add", + "Shanghai", + "2024-03-01T00:00:00Z"); + + function.add(event, accumulator); + List expected = + function.getResult(accumulator); + + function.add(duplicate, accumulator); + Assertions.assertEquals( + expected, + function.getResult(accumulator)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> function.add(conflict, accumulator)); + Assertions.assertEquals( + expected, + function.getResult(accumulator)); + } + + @Test + public void testMergeReplaysEventsWithoutMutatingInputs() { + MemoryEvent add = addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z"); + MemoryEvent correction = correctEvent( + "event-correct", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + MemoryEvent retract = retractEvent( + "event-retract", + "2024-08-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-11-01T00:00:00Z"); + MemoryEvent lateCorrection = correctEvent( + "event-late", + "Tianjin", + "2024-02-01T00:00:00Z", + "2024-03-01T00:00:00Z", + "2024-05-01T00:00:00Z"); + IncrementalTemporalIntegrator left = + function.createAccumulator(); + IncrementalTemporalIntegrator right = + function.createAccumulator(); + addAll(left, add, correction, retract); + addAll(right, add, lateCorrection); + List leftBefore = + function.getResult(left); + List rightBefore = + function.getResult(right); + + IncrementalTemporalIntegrator merged = + function.merge(left, right); + IncrementalTemporalIntegrator mergedAgain = + function.merge(right, left); + + Assertions.assertNotSame(left, merged); + Assertions.assertNotSame(right, merged); + Assertions.assertEquals( + oracle.replay(Arrays.asList( + add, + correction, + retract, + lateCorrection)), + function.getResult(merged)); + Assertions.assertEquals( + function.getResult(merged), + function.getResult(mergedAgain)); + Assertions.assertEquals(leftBefore, function.getResult(left)); + Assertions.assertEquals(rightBefore, function.getResult(right)); + } + + @Test + public void testNullArgumentsRejected() { + IncrementalTemporalIntegrator accumulator = + function.createAccumulator(); + MemoryEvent event = addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z"); + + Assertions.assertThrows( + NullPointerException.class, + () -> function.add(null, accumulator)); + Assertions.assertThrows( + NullPointerException.class, + () -> function.add(event, null)); + Assertions.assertThrows( + NullPointerException.class, + () -> function.getResult(null)); + Assertions.assertThrows( + NullPointerException.class, + () -> function.merge(null, accumulator)); + Assertions.assertThrows( + NullPointerException.class, + () -> function.merge(accumulator, null)); + } + + private void addAll( + IncrementalTemporalIntegrator accumulator, + MemoryEvent... events) { + for (MemoryEvent event : events) { + function.add(event, accumulator); + } + } + + private static MemoryEvent addEvent( + String eventId, + String value, + String transactionTime) { + return MemoryEvent.add( + eventId, + fact(value), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent correctEvent( + String eventId, + String value, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.correct( + eventId, + fact(value), + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent retractEvent( + String eventId, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.retract( + eventId, + "fact-alice-city", + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryFact fact(String value) { + return MemoryFact.attribute( + "fact-alice-city", + new MemoryEntity("person:alice", "person"), + "city", + value); + } + + private static TimeInterval interval( + String start, + String end) { + return new TimeInterval(time(start), time(end)); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "customer-database"), + "Evidence for " + eventId)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalEventPipelineTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalEventPipelineTest.java new file mode 100644 index 000000000..09b36da35 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalEventPipelineTest.java @@ -0,0 +1,389 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.nio.file.StandardOpenOption; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.stream.Stream; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.oracle.FullReplayOracle; +import org.apache.geaflow.api.function.base.KeySelector; +import org.apache.geaflow.api.function.base.MapFunction; +import org.apache.geaflow.api.function.internal.CollectionSource; +import org.apache.geaflow.api.function.io.SinkFunction; +import org.apache.geaflow.api.pdata.stream.window.PWindowSource; +import org.apache.geaflow.api.window.impl.SizeTumblingWindow; +import org.apache.geaflow.cluster.system.ClusterMetaStore; +import org.apache.geaflow.common.config.keys.ExecutionConfigKeys; +import org.apache.geaflow.common.config.keys.FrameworkConfigKeys; +import org.apache.geaflow.env.Environment; +import org.apache.geaflow.env.EnvironmentFactory; +import org.apache.geaflow.file.FileConfigKeys; +import org.apache.geaflow.pipeline.IPipelineResult; +import org.apache.geaflow.pipeline.Pipeline; +import org.apache.geaflow.pipeline.PipelineFactory; +import org.apache.geaflow.pipeline.task.IPipelineTaskContext; +import org.apache.geaflow.pipeline.task.PipelineTask; +import org.apache.geaflow.runtime.core.scheduler.resource.ScheduledWorkerManagerFactory; +import org.apache.geaflow.state.StoreType; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.condition.DisabledOnOs; +import org.junit.jupiter.api.condition.OS; +import org.junit.jupiter.api.io.TempDir; + +public class TemporalEventPipelineTest { + + private static final int WINDOW_SIZE = 2; + + @TempDir + Path tempDirectory; + + @Test + public void testKeyedIncrementalAggregationAcrossWindows() + throws Exception { + assertPipelineResults( + Collections.emptyMap(), + "temporal-results.txt"); + } + + @Test + @DisabledOnOs( + value = OS.WINDOWS, + disabledReason = "GeaFlow LOCAL persistence requires Hadoop winutils.exe") + public void testKeyedAggregationCreatesRocksdbCheckpoints() + throws Exception { + Path checkpointRoot = tempDirectory.resolve("checkpoints"); + assertPipelineResults( + checkpointConfiguration(checkpointRoot), + "temporal-checkpoint-results.txt"); + + Assertions.assertTrue(Files.exists(checkpointRoot)); + try (Stream paths = Files.walk(checkpointRoot)) { + Assertions.assertTrue(paths.anyMatch(path -> + "_commit".equals(path.getFileName().toString()))); + } + } + + private void assertPipelineResults( + Map config, + String outputFileName) throws Exception { + List events = pipelineEvents(); + Path output = tempDirectory.resolve(outputFileName); + Environment environment = null; + + try { + environment = EnvironmentFactory.onLocalEnvironment(); + environment.getEnvironmentContext().withConfig(config); + Pipeline pipeline = + PipelineFactory.buildPipeline(environment); + pipeline.submit(new TemporalPipelineTask( + events, + output.toString())); + + IPipelineResult result = pipeline.execute(); + result.get(); + + Assertions.assertTrue(result.isSuccess()); + List actual = Files.readAllLines( + output, + StandardCharsets.UTF_8); + Collections.sort(actual); + Assertions.assertEquals(5, actual.size()); + Assertions.assertEquals( + expectedWindowResults(events), + actual); + } finally { + if (environment != null) { + environment.shutdown(); + } + ClusterMetaStore.close(); + ScheduledWorkerManagerFactory.clear(); + } + } + + private Map checkpointConfiguration( + Path checkpointRoot) { + Map config = new HashMap<>(); + config.put( + FrameworkConfigKeys.SYSTEM_STATE_BACKEND_TYPE.getKey(), + StoreType.ROCKSDB.name()); + config.put( + FrameworkConfigKeys.BATCH_NUMBER_PER_CHECKPOINT.getKey(), + "1"); + config.put( + ExecutionConfigKeys.JOB_APP_NAME.getKey(), + "TemporalEventPipelineCheckpointTest"); + config.put( + ExecutionConfigKeys.JOB_WORK_PATH.getKey(), + tempDirectory.resolve("work").toString()); + config.put( + FileConfigKeys.PERSISTENT_TYPE.getKey(), + "LOCAL"); + config.put( + FileConfigKeys.ROOT.getKey(), + checkpointRoot.toString()); + return config; + } + + private static List pipelineEvents() { + MemoryEvent bobAdd = addEvent( + "event-bob-add", + "fact-bob-city", + "person:bob", + "Paris", + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"); + List events = new ArrayList<>(); + events.add(addEvent( + "event-alice-add", + "fact-alice-city", + "person:alice", + "Beijing", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z")); + events.add(bobAdd); + events.add(correctEvent( + "event-alice-correct", + "fact-alice-city", + "person:alice", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z")); + events.add(bobAdd); + events.add(retractEvent( + "event-alice-retract", + "fact-alice-city", + "2024-08-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-11-01T00:00:00Z")); + events.add(correctEvent( + "event-alice-late", + "fact-alice-city", + "person:alice", + "Tianjin", + "2024-02-01T00:00:00Z", + "2024-03-01T00:00:00Z", + "2024-05-01T00:00:00Z")); + return events; + } + + private static List expectedWindowResults( + List events) { + FullReplayOracle oracle = new FullReplayOracle(); + Map> receivedByFact = + new HashMap<>(); + List expected = new ArrayList<>(); + + for (int start = 0; + start < events.size(); + start += WINDOW_SIZE) { + Set touchedFactIds = new LinkedHashSet<>(); + int end = Math.min(start + WINDOW_SIZE, events.size()); + for (int index = start; index < end; index++) { + MemoryEvent event = events.get(index); + receivedByFact.computeIfAbsent( + event.getFactId(), + ignored -> new ArrayList<>()).add(event); + touchedFactIds.add(event.getFactId()); + } + for (String factId : touchedFactIds) { + expected.add(formatSnapshot( + oracle.replay(receivedByFact.get(factId)))); + } + } + + Collections.sort(expected); + return expected; + } + + private static String formatSnapshot( + List versions) { + StringBuilder builder = new StringBuilder(); + for (MemoryFactVersion version : versions) { + if (builder.length() > 0) { + builder.append(';'); + } + builder.append(version.getId()) + .append('|') + .append(version.getFact().getId()) + .append('|') + .append(version.getFact().getLiteralValue().get()) + .append('|') + .append(version.getValidTime()) + .append('|') + .append(version.getTransactionTime()); + } + return builder.toString(); + } + + private static MemoryEvent addEvent( + String eventId, + String factId, + String subjectId, + String value, + String validStart, + String transactionTime) { + return MemoryEvent.add( + eventId, + fact(factId, subjectId, value), + TimeInterval.unboundedFrom(time(validStart)), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent correctEvent( + String eventId, + String factId, + String subjectId, + String value, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.correct( + eventId, + fact(factId, subjectId, value), + new TimeInterval(time(validStart), time(validEnd)), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent retractEvent( + String eventId, + String factId, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.retract( + eventId, + factId, + new TimeInterval(time(validStart), time(validEnd)), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryFact fact( + String factId, + String subjectId, + String value) { + return MemoryFact.attribute( + factId, + new MemoryEntity(subjectId, "person"), + "city", + value); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "customer-database"), + "Evidence for " + eventId)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } + + private static final class TemporalPipelineTask + implements PipelineTask { + + private final List events; + private final String outputPath; + + private TemporalPipelineTask( + List events, + String outputPath) { + this.events = new ArrayList<>(events); + this.outputPath = outputPath; + } + + @Override + public void execute( + IPipelineTaskContext pipelineTaskContext) { + PWindowSource source = + pipelineTaskContext.buildSource( + new CollectionSource<>(events), + SizeTumblingWindow.of(WINDOW_SIZE)); + source.withParallelism(1) + .keyBy(new FactIdSelector()) + .aggregate(new TemporalEventAggregateFunction()) + .withParallelism(2) + .map(new SnapshotFormatter()) + .sink(new LineFileSink(outputPath)) + .withParallelism(1); + } + } + + private static final class FactIdSelector implements + KeySelector { + + @Override + public String getKey(MemoryEvent event) { + return event.getFactId(); + } + } + + private static final class SnapshotFormatter implements + MapFunction, String> { + + @Override + public String map(List versions) { + return formatSnapshot(versions); + } + } + + private static final class LineFileSink implements + SinkFunction { + + private final String outputPath; + + private LineFileSink(String outputPath) { + this.outputPath = outputPath; + } + + @Override + public void write(String value) throws Exception { + Files.write( + Paths.get(outputPath), + Collections.singletonList(value), + StandardCharsets.UTF_8, + StandardOpenOption.CREATE, + StandardOpenOption.APPEND); + } + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalGeaFlowSerializationTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalGeaFlowSerializationTest.java new file mode 100644 index 000000000..5e6014987 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalGeaFlowSerializationTest.java @@ -0,0 +1,210 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.time.Instant; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.oracle.FullReplayOracle; +import org.apache.geaflow.common.serialize.ISerializer; +import org.apache.geaflow.common.serialize.SerializerFactory; +import org.apache.geaflow.state.serializer.DefaultKVSerializer; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class TemporalGeaFlowSerializationTest { + + private final TemporalEventAggregateFunction function = + new TemporalEventAggregateFunction(); + private final FullReplayOracle oracle = new FullReplayOracle(); + + @Test + public void testMemoryEventRoundTripsThroughShuffleSerializer() { + MemoryEvent event = correctEvent( + "event-correct", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + ISerializer serializer = + SerializerFactory.getKryoSerializer(); + + MemoryEvent restored = (MemoryEvent) serializer.deserialize( + serializer.serialize(event)); + + Assertions.assertEquals(event, restored); + } + + @Test + public void testAccumulatorRoundTripsThroughKeyValueStateSerializer() { + MemoryEvent add = addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z"); + MemoryEvent correction = correctEvent( + "event-correct", + "Shanghai", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + MemoryEvent retract = retractEvent( + "event-retract", + "2024-08-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-11-01T00:00:00Z"); + MemoryEvent lateCorrection = correctEvent( + "event-late", + "Tianjin", + "2024-02-01T00:00:00Z", + "2024-03-01T00:00:00Z", + "2024-05-01T00:00:00Z"); + IncrementalTemporalIntegrator accumulator = + function.createAccumulator(); + function.add(add, accumulator); + function.add(correction, accumulator); + function.add(retract, accumulator); + DefaultKVSerializer + serializer = new DefaultKVSerializer<>(String.class, null); + + Assertions.assertEquals( + "fact-alice-city", + serializer.deserializeKey( + serializer.serializeKey("fact-alice-city"))); + IncrementalTemporalIntegrator restored = + serializer.deserializeValue( + serializer.serializeValue(accumulator)); + + Assertions.assertNotNull(restored); + Assertions.assertEquals( + accumulator.eventSnapshot(), + restored.eventSnapshot()); + Assertions.assertEquals( + function.getResult(accumulator), + function.getResult(restored)); + + function.add(lateCorrection, restored); + Assertions.assertEquals( + oracle.replay(Arrays.asList( + add, + correction, + retract, + lateCorrection)), + function.getResult(restored)); + } + + @Test + @SuppressWarnings("unchecked") + public void testResultSnapshotKeepsImmutabilityAfterRoundTrip() { + IncrementalTemporalIntegrator accumulator = + function.createAccumulator(); + function.add( + addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z"), + accumulator); + List snapshot = + function.getResult(accumulator); + ISerializer serializer = + SerializerFactory.getKryoSerializer(); + + List restored = + (List) serializer.deserialize( + serializer.serialize(snapshot)); + + Assertions.assertEquals(snapshot, restored); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> restored.clear()); + } + + private static MemoryEvent addEvent( + String eventId, + String value, + String transactionTime) { + return MemoryEvent.add( + eventId, + fact(value), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent correctEvent( + String eventId, + String value, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.correct( + eventId, + fact(value), + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent retractEvent( + String eventId, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.retract( + eventId, + "fact-alice-city", + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryFact fact(String value) { + return MemoryFact.attribute( + "fact-alice-city", + new MemoryEntity("person:alice", "person"), + "city", + value); + } + + private static TimeInterval interval( + String start, + String end) { + return new TimeInterval(time(start), time(end)); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "customer-database"), + "Evidence for " + eventId)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalStateRecoveryTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalStateRecoveryTest.java new file mode 100644 index 000000000..b1cdab05c --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/integration/TemporalStateRecoveryTest.java @@ -0,0 +1,636 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.integration; + +import java.nio.file.Path; +import java.time.Instant; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.oracle.FullReplayOracle; +import org.apache.geaflow.common.config.Configuration; +import org.apache.geaflow.common.config.keys.ExecutionConfigKeys; +import org.apache.geaflow.file.FileConfigKeys; +import org.apache.geaflow.state.KeyValueState; +import org.apache.geaflow.state.StateFactory; +import org.apache.geaflow.state.StoreType; +import org.apache.geaflow.state.descriptor.KeyValueStateDescriptor; +import org.apache.geaflow.utils.keygroup.DefaultKeyGroupAssigner; +import org.apache.geaflow.utils.keygroup.KeyGroup; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.condition.DisabledOnOs; +import org.junit.jupiter.api.condition.OS; +import org.junit.jupiter.api.io.TempDir; + +@DisabledOnOs( + value = OS.WINDOWS, + disabledReason = "GeaFlow LOCAL persistence requires Hadoop winutils.exe") +public class TemporalStateRecoveryTest { + + private static final String FACT_ID = "fact-alice-city"; + private static final String BOB_FACT_ID = "fact-bob-city"; + private static final long CHECKPOINT_ID = 1L; + + private final FullReplayOracle oracle = new FullReplayOracle(); + + @TempDir + Path tempDirectory; + + @Test + public void testAccumulatorRecoversFromRocksdbCheckpoint() { + MemoryEvent add = addEvent(); + IncrementalTemporalIntegrator accumulator = + new IncrementalTemporalIntegrator(); + accumulator.apply(add); + + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState originalState = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + originalState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + originalState.put(FACT_ID, accumulator); + originalState.manage().operate().finish(); + originalState.manage().operate().archive(); + } finally { + closeAndDrop(originalState); + } + + KeyValueState recoveredState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + recoveredState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + recoveredState.manage().operate().recover(); + + IncrementalTemporalIntegrator recovered = + recoveredState.get(FACT_ID); + Assertions.assertNotNull(recovered); + Assertions.assertEquals( + Collections.singletonList(add), + recovered.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Collections.singletonList(add)), + recovered.snapshot()); + } finally { + closeAndDrop(recoveredState); + } + } + + @Test + public void testRecoverDiscardsUncheckpointedChanges() { + MemoryEvent add = addEvent(); + MemoryEvent correction = correctEvent(); + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState state = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + IncrementalTemporalIntegrator accumulator = + new IncrementalTemporalIntegrator(); + accumulator.apply(add); + state.manage().operate().setCheckpointId(CHECKPOINT_ID); + state.put(FACT_ID, accumulator); + state.manage().operate().finish(); + state.manage().operate().archive(); + + state.manage().operate() + .setCheckpointId(CHECKPOINT_ID + 1); + IncrementalTemporalIntegrator uncheckpointed = + state.get(FACT_ID); + uncheckpointed.apply(correction); + state.put(FACT_ID, uncheckpointed); + Assertions.assertEquals( + oracle.replay(Arrays.asList(add, correction)), + state.get(FACT_ID).snapshot()); + + state.manage().operate().setCheckpointId(CHECKPOINT_ID); + state.manage().operate().recover(); + + IncrementalTemporalIntegrator recovered = + state.get(FACT_ID); + Assertions.assertNotNull(recovered); + Assertions.assertEquals( + Collections.singletonList(add), + recovered.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Collections.singletonList(add)), + recovered.snapshot()); + } finally { + closeAndDrop(state); + } + } + + @Test + public void testRecoveredAccumulatorHandlesLateCorrection() { + MemoryEvent add = addEvent(); + MemoryEvent correction = correctEvent(); + MemoryEvent lateCorrection = lateCorrectionEvent(); + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState originalState = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + IncrementalTemporalIntegrator accumulator = + new IncrementalTemporalIntegrator(); + accumulator.apply(add); + accumulator.apply(correction); + originalState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + originalState.put(FACT_ID, accumulator); + originalState.manage().operate().finish(); + originalState.manage().operate().archive(); + } finally { + closeAndDrop(originalState); + } + + KeyValueState recoveredState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + recoveredState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + recoveredState.manage().operate().recover(); + + IncrementalTemporalIntegrator recovered = + recoveredState.get(FACT_ID); + Assertions.assertNotNull(recovered); + recovered.apply(lateCorrection); + recoveredState.put(FACT_ID, recovered); + + IncrementalTemporalIntegrator updated = + recoveredState.get(FACT_ID); + Assertions.assertEquals( + Arrays.asList(add, lateCorrection, correction), + updated.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Arrays.asList( + add, + correction, + lateCorrection)), + updated.snapshot()); + } finally { + closeAndDrop(recoveredState); + } + } + + @Test + public void testContinuedUpdatesSurviveNextCheckpoint() { + MemoryEvent add = addEvent(); + MemoryEvent correction = correctEvent(); + MemoryEvent lateCorrection = lateCorrectionEvent(); + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState originalState = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + IncrementalTemporalIntegrator accumulator = + new IncrementalTemporalIntegrator(); + accumulator.apply(add); + accumulator.apply(correction); + originalState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + originalState.put(FACT_ID, accumulator); + originalState.manage().operate().finish(); + originalState.manage().operate().archive(); + } finally { + closeAndDrop(originalState); + } + + KeyValueState continuedState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + continuedState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + continuedState.manage().operate().recover(); + + IncrementalTemporalIntegrator recovered = + continuedState.get(FACT_ID); + Assertions.assertNotNull(recovered); + recovered.apply(lateCorrection); + continuedState.put(FACT_ID, recovered); + continuedState.manage().operate() + .setCheckpointId(CHECKPOINT_ID + 1); + continuedState.manage().operate().finish(); + continuedState.manage().operate().archive(); + } finally { + closeAndDrop(continuedState); + } + + KeyValueState restoredState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + restoredState.manage().operate() + .setCheckpointId(CHECKPOINT_ID + 1); + restoredState.manage().operate().recover(); + + IncrementalTemporalIntegrator restored = + restoredState.get(FACT_ID); + Assertions.assertNotNull(restored); + Assertions.assertEquals( + Arrays.asList(add, lateCorrection, correction), + restored.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Arrays.asList( + add, + correction, + lateCorrection)), + restored.snapshot()); + } finally { + closeAndDrop(restoredState); + } + } + + @Test + public void testConflictingEventAfterRecoveryIsAtomic() { + MemoryEvent add = addEvent(); + MemoryEvent conflictingAdd = conflictingAddEvent(); + MemoryEvent correction = correctEvent(); + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState originalState = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + IncrementalTemporalIntegrator accumulator = + new IncrementalTemporalIntegrator(); + accumulator.apply(add); + originalState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + originalState.put(FACT_ID, accumulator); + originalState.manage().operate().finish(); + originalState.manage().operate().archive(); + } finally { + closeAndDrop(originalState); + } + + KeyValueState recoveredState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + recoveredState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + recoveredState.manage().operate().recover(); + + IncrementalTemporalIntegrator recovered = + recoveredState.get(FACT_ID); + Assertions.assertNotNull(recovered); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> recovered.apply(conflictingAdd)); + Assertions.assertEquals( + Collections.singletonList(add), + recovered.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Collections.singletonList(add)), + recovered.snapshot()); + + recovered.apply(correction); + recoveredState.put(FACT_ID, recovered); + IncrementalTemporalIntegrator updated = + recoveredState.get(FACT_ID); + Assertions.assertEquals( + Arrays.asList(add, correction), + updated.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Arrays.asList(add, correction)), + updated.snapshot()); + } finally { + closeAndDrop(recoveredState); + } + } + + @Test + public void testMultipleFactIdsRecoverIndependently() { + MemoryEvent aliceAdd = addEvent(); + MemoryEvent aliceCorrection = correctEvent(); + MemoryEvent bobAdd = bobAddEvent(); + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState originalState = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + IncrementalTemporalIntegrator aliceAccumulator = + new IncrementalTemporalIntegrator(); + aliceAccumulator.apply(aliceAdd); + IncrementalTemporalIntegrator bobAccumulator = + new IncrementalTemporalIntegrator(); + bobAccumulator.apply(bobAdd); + + originalState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + originalState.put(FACT_ID, aliceAccumulator); + originalState.put(BOB_FACT_ID, bobAccumulator); + originalState.manage().operate().finish(); + originalState.manage().operate().archive(); + } finally { + closeAndDrop(originalState); + } + + KeyValueState recoveredState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + recoveredState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + recoveredState.manage().operate().recover(); + + IncrementalTemporalIntegrator recoveredAlice = + recoveredState.get(FACT_ID); + IncrementalTemporalIntegrator recoveredBob = + recoveredState.get(BOB_FACT_ID); + Assertions.assertNotNull(recoveredAlice); + Assertions.assertNotNull(recoveredBob); + Assertions.assertEquals( + Collections.singletonList(aliceAdd), + recoveredAlice.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Collections.singletonList(aliceAdd)), + recoveredAlice.snapshot()); + Assertions.assertEquals( + Collections.singletonList(bobAdd), + recoveredBob.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Collections.singletonList(bobAdd)), + recoveredBob.snapshot()); + + recoveredAlice.apply(aliceCorrection); + recoveredState.put(FACT_ID, recoveredAlice); + + IncrementalTemporalIntegrator updatedAlice = + recoveredState.get(FACT_ID); + IncrementalTemporalIntegrator unchangedBob = + recoveredState.get(BOB_FACT_ID); + Assertions.assertEquals( + Arrays.asList(aliceAdd, aliceCorrection), + updatedAlice.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Arrays.asList( + aliceAdd, + aliceCorrection)), + updatedAlice.snapshot()); + Assertions.assertEquals( + Collections.singletonList(bobAdd), + unchangedBob.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Collections.singletonList(bobAdd)), + unchangedBob.snapshot()); + } finally { + closeAndDrop(recoveredState); + } + } + + @Test + public void testDuplicateEventAfterRecoveryIsIdempotent() { + MemoryEvent add = addEvent(); + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState originalState = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + IncrementalTemporalIntegrator accumulator = + new IncrementalTemporalIntegrator(); + accumulator.apply(add); + originalState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + originalState.put(FACT_ID, accumulator); + originalState.manage().operate().finish(); + originalState.manage().operate().archive(); + } finally { + closeAndDrop(originalState); + } + + KeyValueState recoveredState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + recoveredState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + recoveredState.manage().operate().recover(); + + IncrementalTemporalIntegrator recovered = + recoveredState.get(FACT_ID); + Assertions.assertNotNull(recovered); + recovered.apply(add); + recoveredState.put(FACT_ID, recovered); + + IncrementalTemporalIntegrator updated = + recoveredState.get(FACT_ID); + Assertions.assertEquals( + Collections.singletonList(add), + updated.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Collections.singletonList(add)), + updated.snapshot()); + } finally { + closeAndDrop(recoveredState); + } + } + + @Test + public void testRecoveredAccumulatorHandlesRetraction() { + MemoryEvent add = addEvent(); + MemoryEvent correction = correctEvent(); + MemoryEvent retraction = retractEvent(); + Configuration configuration = stateConfiguration(); + KeyValueStateDescriptor + descriptor = stateDescriptor(); + KeyValueState originalState = + StateFactory.buildKeyValueState(descriptor, configuration); + + try { + IncrementalTemporalIntegrator accumulator = + new IncrementalTemporalIntegrator(); + accumulator.apply(add); + accumulator.apply(correction); + originalState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + originalState.put(FACT_ID, accumulator); + originalState.manage().operate().finish(); + originalState.manage().operate().archive(); + } finally { + closeAndDrop(originalState); + } + + KeyValueState recoveredState = + StateFactory.buildKeyValueState(descriptor, configuration); + try { + recoveredState.manage().operate() + .setCheckpointId(CHECKPOINT_ID); + recoveredState.manage().operate().recover(); + + IncrementalTemporalIntegrator recovered = + recoveredState.get(FACT_ID); + Assertions.assertNotNull(recovered); + recovered.apply(retraction); + recoveredState.put(FACT_ID, recovered); + + IncrementalTemporalIntegrator updated = + recoveredState.get(FACT_ID); + Assertions.assertEquals( + Arrays.asList(add, correction, retraction), + updated.eventSnapshot()); + Assertions.assertEquals( + oracle.replay(Arrays.asList( + add, + correction, + retraction)), + updated.snapshot()); + } finally { + closeAndDrop(recoveredState); + } + } + + private Configuration stateConfiguration() { + Map config = new HashMap<>(); + config.put( + ExecutionConfigKeys.JOB_APP_NAME.getKey(), + "TemporalStateRecoveryTest"); + config.put( + ExecutionConfigKeys.JOB_WORK_PATH.getKey(), + tempDirectory.resolve("work").toString()); + config.put( + FileConfigKeys.PERSISTENT_TYPE.getKey(), + "LOCAL"); + config.put( + FileConfigKeys.ROOT.getKey(), + tempDirectory.resolve("checkpoints").toString()); + return new Configuration(config); + } + + private static KeyValueStateDescriptor stateDescriptor() { + KeyValueStateDescriptor + descriptor = KeyValueStateDescriptor.build( + "temporal-recovery", + StoreType.ROCKSDB.name()); + descriptor.withKeyGroup(new KeyGroup(0, 0)) + .withKeyGroupAssigner(new DefaultKeyGroupAssigner(1)); + return descriptor; + } + + private static void closeAndDrop( + KeyValueState state) { + state.manage().operate().close(); + state.manage().operate().drop(); + } + + private static MemoryEvent addEvent() { + return addEvent("Beijing"); + } + + private static MemoryEvent conflictingAddEvent() { + return addEvent("Shenzhen"); + } + + private static MemoryEvent bobAddEvent() { + return MemoryEvent.add( + "event-bob-add", + fact(BOB_FACT_ID, "person:bob", "Paris"), + TimeInterval.unboundedFrom( + Instant.parse("2024-01-01T00:00:00Z")), + Instant.parse("2024-04-01T00:00:00Z"), + evidence("event-bob-add")); + } + + private static MemoryEvent addEvent(String value) { + return MemoryEvent.add( + "event-alice-add", + fact(value), + TimeInterval.unboundedFrom( + Instant.parse("2024-01-01T00:00:00Z")), + Instant.parse("2024-03-01T00:00:00Z"), + evidence("event-alice-add")); + } + + private static MemoryEvent correctEvent() { + return MemoryEvent.correct( + "event-alice-correct", + fact("Shanghai"), + new TimeInterval( + Instant.parse("2024-04-01T00:00:00Z"), + Instant.parse("2024-09-01T00:00:00Z")), + Instant.parse("2024-06-01T00:00:00Z"), + evidence("event-alice-correct")); + } + + private static MemoryEvent lateCorrectionEvent() { + return MemoryEvent.correct( + "event-alice-late", + fact("Tianjin"), + new TimeInterval( + Instant.parse("2024-02-01T00:00:00Z"), + Instant.parse("2024-03-01T00:00:00Z")), + Instant.parse("2024-05-01T00:00:00Z"), + evidence("event-alice-late")); + } + + private static MemoryEvent retractEvent() { + return MemoryEvent.retract( + "event-alice-retract", + FACT_ID, + new TimeInterval( + Instant.parse("2024-08-01T00:00:00Z"), + Instant.parse("2024-10-01T00:00:00Z")), + Instant.parse("2024-11-01T00:00:00Z"), + evidence("event-alice-retract")); + } + + private static MemoryFact fact(String value) { + return fact(FACT_ID, "person:alice", value); + } + + private static MemoryFact fact( + String factId, + String subjectId, + String value) { + return MemoryFact.attribute( + factId, + new MemoryEntity(subjectId, "person"), + "city", + value); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "customer-database"), + "Evidence for " + eventId)); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/FactKeyTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/FactKeyTest.java new file mode 100644 index 000000000..b90015ca0 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/FactKeyTest.java @@ -0,0 +1,98 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class FactKeyTest { + + @Test + public void testValueSemanticsAndDeterministicOrder() { + FactKey key = new FactKey( + "person:alice", + "city", + "profile"); + FactKey same = new FactKey( + "person:alice", + "city", + "profile"); + + Assertions.assertEquals("person:alice", key.getSubjectId()); + Assertions.assertEquals("city", key.getPredicate()); + Assertions.assertEquals("profile", key.getScope()); + Assertions.assertEquals(key, same); + Assertions.assertEquals(key.hashCode(), same.hashCode()); + + FactKey earlierPredicate = new FactKey( + "person:alice", + "age", + "profile"); + FactKey earlierScope = new FactKey( + "person:alice", + "city", + "account"); + FactKey laterSubject = new FactKey( + "person:bob", + "age", + "profile"); + List keys = new ArrayList<>(Arrays.asList( + laterSubject, + key, + earlierScope, + earlierPredicate)); + + Collections.sort(keys); + + Assertions.assertEquals( + Arrays.asList( + earlierPredicate, + earlierScope, + key, + laterSubject), + keys); + } + + @Test + public void testRejectInvalidKey() { + Assertions.assertThrows( + NullPointerException.class, + () -> new FactKey(null, "city", "profile")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new FactKey(" ", "city", "profile")); + Assertions.assertThrows( + NullPointerException.class, + () -> new FactKey("person:alice", null, "profile")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new FactKey("person:alice", " ", "profile")); + Assertions.assertThrows( + NullPointerException.class, + () -> new FactKey("person:alice", "city", null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new FactKey("person:alice", "city", " ")); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/FactValueTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/FactValueTest.java new file mode 100644 index 000000000..25d215a80 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/FactValueTest.java @@ -0,0 +1,101 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class FactValueTest { + + @Test + public void testLiteralAndEntityReference() { + FactValue literal = FactValue.literal("Beijing"); + FactValue entityReference = + FactValue.entityReference("city:beijing"); + + Assertions.assertEquals( + FactValue.Kind.LITERAL, + literal.getKind()); + Assertions.assertEquals("Beijing", literal.getValue()); + Assertions.assertEquals( + Optional.of("Beijing"), + literal.getLiteralValue()); + Assertions.assertEquals( + Optional.empty(), + literal.getEntityId()); + + Assertions.assertEquals( + FactValue.Kind.ENTITY_REF, + entityReference.getKind()); + Assertions.assertEquals( + "city:beijing", + entityReference.getValue()); + Assertions.assertEquals( + Optional.empty(), + entityReference.getLiteralValue()); + Assertions.assertEquals( + Optional.of("city:beijing"), + entityReference.getEntityId()); + } + + @Test + public void testValueSemanticsAndDeterministicOrder() { + FactValue beijing = FactValue.literal("Beijing"); + FactValue same = FactValue.literal("Beijing"); + FactValue shanghai = FactValue.literal("Shanghai"); + FactValue reference = + FactValue.entityReference("city:beijing"); + + Assertions.assertEquals(beijing, same); + Assertions.assertEquals(beijing.hashCode(), same.hashCode()); + Assertions.assertNotEquals(beijing, reference); + + List values = new ArrayList<>(Arrays.asList( + reference, + shanghai, + beijing)); + Collections.sort(values); + + Assertions.assertEquals( + Arrays.asList(beijing, shanghai, reference), + values); + } + + @Test + public void testRejectInvalidValue() { + Assertions.assertThrows( + NullPointerException.class, + () -> FactValue.literal(null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> FactValue.literal(" ")); + Assertions.assertThrows( + NullPointerException.class, + () -> FactValue.entityReference(null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> FactValue.entityReference(" ")); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryEntityTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryEntityTest.java new file mode 100644 index 000000000..61a54fb7c --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryEntityTest.java @@ -0,0 +1,62 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class MemoryEntityTest { + + @Test + public void testValueSemantics() { + MemoryEntity entity = + new MemoryEntity("person:alice", "person"); + MemoryEntity same = + new MemoryEntity("person:alice", "person"); + + Assertions.assertEquals("person:alice", entity.getId()); + Assertions.assertEquals("person", entity.getLabel()); + Assertions.assertEquals(entity, same); + Assertions.assertEquals(entity.hashCode(), same.hashCode()); + + Assertions.assertNotEquals( + entity, + new MemoryEntity("person:bob", "person")); + Assertions.assertNotEquals( + entity, + new MemoryEntity("person:alice", "company")); + } + + @Test + public void testRejectInvalidEntity() { + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryEntity(null, "person")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new MemoryEntity(" ", "person")); + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryEntity("person:alice", null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new MemoryEntity("person:alice", " ")); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryEventTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryEventTest.java new file mode 100644 index 000000000..7c4c32967 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryEventTest.java @@ -0,0 +1,272 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class MemoryEventTest { + + @Test + public void testAddLateEvent() { + MemoryFact fact = fact("Alice"); + TimeInterval validTime = TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")); + Instant transactionTime = + time("2024-03-01T00:00:00Z"); + List evidence = evidenceList(); + + MemoryEvent event = MemoryEvent.add( + "event-1", + fact, + validTime, + transactionTime, + evidence); + + Assertions.assertEquals("event-1", event.getId()); + Assertions.assertEquals( + MemoryEventOperation.ADD, + event.getOperation()); + Assertions.assertEquals("fact-name-alice", event.getFactId()); + Assertions.assertEquals(fact, event.getFact().get()); + Assertions.assertEquals(validTime, event.getValidTime()); + Assertions.assertEquals( + transactionTime, + event.getTransactionTime()); + Assertions.assertEquals(evidence, event.getEvidence()); + Assertions.assertTrue( + event.getValidTime().getStart() + .isBefore(event.getTransactionTime())); + } + + @Test + public void testCorrectEvent() { + MemoryFact corrected = fact("Alice Smith"); + + MemoryEvent event = MemoryEvent.correct( + "event-2", + corrected, + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time("2024-06-01T00:00:00Z"), + evidenceList()); + + Assertions.assertEquals( + MemoryEventOperation.CORRECT, + event.getOperation()); + Assertions.assertEquals("fact-name-alice", event.getFactId()); + Assertions.assertEquals(corrected, event.getFact().get()); + } + + @Test + public void testRetractEvent() { + TimeInterval validTime = TimeInterval.unboundedFrom( + time("2025-01-01T00:00:00Z")); + + MemoryEvent event = MemoryEvent.retract( + "event-3", + "fact-name-alice", + validTime, + time("2025-02-01T00:00:00Z"), + evidenceList()); + + Assertions.assertEquals( + MemoryEventOperation.RETRACT, + event.getOperation()); + Assertions.assertEquals("fact-name-alice", event.getFactId()); + Assertions.assertFalse(event.getFact().isPresent()); + Assertions.assertEquals(validTime, event.getValidTime()); + } + + @Test + public void testValueSemanticsAndEvidenceCopy() { + Evidence first = evidence( + "evidence-1", + "Alice is the recorded name"); + List mutableEvidence = new ArrayList<>(); + mutableEvidence.add(first); + + MemoryEvent event = MemoryEvent.add( + "event-1", + fact("Alice"), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time("2024-03-01T00:00:00Z"), + mutableEvidence); + + MemoryEvent same = MemoryEvent.add( + "event-1", + fact("Alice"), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time("2024-03-01T00:00:00Z"), + Collections.singletonList(evidence( + "evidence-1", + "Alice is the recorded name"))); + + mutableEvidence.add(evidence( + "evidence-2", + "An independent record")); + + Assertions.assertEquals( + Collections.singletonList(first), + event.getEvidence()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> event.getEvidence().clear()); + Assertions.assertEquals(event, same); + Assertions.assertEquals(event.hashCode(), same.hashCode()); + + Assertions.assertNotEquals( + event, + MemoryEvent.add( + "event-1", + fact("Alice Smith"), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time("2024-03-01T00:00:00Z"), + evidenceList())); + } + + @Test + public void testRejectInvalidEvent() { + MemoryFact fact = fact("Alice"); + TimeInterval validTime = TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")); + Instant transactionTime = + time("2024-03-01T00:00:00Z"); + List evidence = evidenceList(); + + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryEvent.add( + null, + fact, + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> MemoryEvent.add( + " ", + fact, + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryEvent.add( + "event-1", + null, + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryEvent.retract( + "event-1", + null, + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> MemoryEvent.retract( + "event-1", + " ", + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryEvent.add( + "event-1", + fact, + null, + transactionTime, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryEvent.add( + "event-1", + fact, + validTime, + null, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryEvent.add( + "event-1", + fact, + validTime, + transactionTime, + null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> MemoryEvent.add( + "event-1", + fact, + validTime, + transactionTime, + Collections.emptyList())); + + List evidenceWithNull = new ArrayList<>(); + evidenceWithNull.add(null); + + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryEvent.add( + "event-1", + fact, + validTime, + transactionTime, + evidenceWithNull)); + } + + private static MemoryFact fact(String literalValue) { + return MemoryFact.attribute( + "fact-name-alice", + new MemoryEntity("person:alice", "person"), + "name", + literalValue); + } + + private static List evidenceList() { + return Collections.singletonList(evidence( + "evidence-1", + "Alice is the recorded name")); + } + + private static Evidence evidence(String id, String content) { + return new Evidence( + id, + new Source("source-1", "customer-database"), + content); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryFactTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryFactTest.java new file mode 100644 index 000000000..0e170d372 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryFactTest.java @@ -0,0 +1,141 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class MemoryFactTest { + + @Test + public void testAttributeFact() { + MemoryEntity alice = + new MemoryEntity("person:alice", "person"); + MemoryFact fact = MemoryFact.attribute( + "fact-name-alice", + alice, + "name", + "Alice"); + MemoryFact same = MemoryFact.attribute( + "fact-name-alice", + new MemoryEntity("person:alice", "person"), + "name", + "Alice"); + + Assertions.assertEquals("fact-name-alice", fact.getId()); + Assertions.assertEquals(alice, fact.getSubject()); + Assertions.assertEquals("name", fact.getPredicate()); + Assertions.assertFalse(fact.isRelationship()); + Assertions.assertEquals( + "Alice", + fact.getLiteralValue().get()); + Assertions.assertFalse(fact.getTarget().isPresent()); + Assertions.assertEquals(fact, same); + Assertions.assertEquals(fact.hashCode(), same.hashCode()); + + Assertions.assertNotEquals( + fact, + MemoryFact.attribute( + "fact-name-alice", + alice, + "name", + "Alice Smith")); + } + + @Test + public void testRelationshipFact() { + MemoryEntity alice = + new MemoryEntity("person:alice", "person"); + MemoryEntity acme = + new MemoryEntity("company:acme", "company"); + MemoryFact fact = MemoryFact.relationship( + "fact-alice-acme", + alice, + "worksAt", + acme); + MemoryFact same = MemoryFact.relationship( + "fact-alice-acme", + new MemoryEntity("person:alice", "person"), + "worksAt", + new MemoryEntity("company:acme", "company")); + + Assertions.assertTrue(fact.isRelationship()); + Assertions.assertFalse(fact.getLiteralValue().isPresent()); + Assertions.assertEquals(acme, fact.getTarget().get()); + Assertions.assertEquals(fact, same); + Assertions.assertEquals(fact.hashCode(), same.hashCode()); + + Assertions.assertNotEquals( + fact, + MemoryFact.relationship( + "fact-alice-acme", + alice, + "worksAt", + new MemoryEntity("company:other", "company"))); + } + + @Test + public void testRejectInvalidFact() { + MemoryEntity alice = + new MemoryEntity("person:alice", "person"); + + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryFact.attribute(null, alice, "name", "Alice")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> MemoryFact.attribute(" ", alice, "name", "Alice")); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryFact.attribute( + "fact-name-alice", + null, + "name", + "Alice")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> MemoryFact.attribute( + "fact-name-alice", + alice, + " ", + "Alice")); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryFact.attribute( + "fact-name-alice", + alice, + "name", + null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> MemoryFact.attribute( + "fact-name-alice", + alice, + "name", + " ")); + Assertions.assertThrows( + NullPointerException.class, + () -> MemoryFact.relationship( + "fact-alice-acme", + alice, + "worksAt", + null)); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersionTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersionTest.java new file mode 100644 index 000000000..2f03faa86 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/MemoryFactVersionTest.java @@ -0,0 +1,297 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class MemoryFactVersionTest { + + @Test + public void testBitemporalVersion() { + MemoryFact fact = fact("Alice"); + TimeInterval validTime = TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")); + TimeInterval transactionTime = interval( + "2024-03-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + Evidence evidence = evidence( + "evidence-1", + "Alice is the recorded name"); + + MemoryFactVersion version = new MemoryFactVersion( + "version-1", + fact, + validTime, + transactionTime, + Collections.singletonList(evidence)); + + Assertions.assertEquals("version-1", version.getId()); + Assertions.assertEquals(fact, version.getFact()); + Assertions.assertEquals(validTime, version.getValidTime()); + Assertions.assertEquals( + transactionTime, + version.getTransactionTime()); + Assertions.assertEquals( + Collections.singletonList(evidence), + version.getEvidence()); + } + + @Test + public void testValueSemantics() { + MemoryFactVersion version = new MemoryFactVersion( + "version-1", + fact("Alice"), + interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + TimeInterval.unboundedFrom( + time("2024-03-01T00:00:00Z")), + Collections.singletonList(evidence( + "evidence-1", + "Alice is the recorded name"))); + + MemoryFactVersion same = new MemoryFactVersion( + "version-1", + fact("Alice"), + interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + TimeInterval.unboundedFrom( + time("2024-03-01T00:00:00Z")), + Collections.singletonList(evidence( + "evidence-1", + "Alice is the recorded name"))); + + Assertions.assertEquals(version, same); + Assertions.assertEquals(version.hashCode(), same.hashCode()); + + Assertions.assertNotEquals( + version, + new MemoryFactVersion( + "version-1", + fact("Alice Smith"), + interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + TimeInterval.unboundedFrom( + time("2024-03-01T00:00:00Z")), + Collections.singletonList(evidence( + "evidence-1", + "Alice is the recorded name")))); + + Assertions.assertNotEquals( + version, + new MemoryFactVersion( + "version-1", + fact("Alice"), + interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + TimeInterval.unboundedFrom( + time("2024-04-01T00:00:00Z")), + Collections.singletonList(evidence( + "evidence-1", + "Alice is the recorded name")))); + } + + @Test + public void testEvidenceIsDefensivelyCopied() { + Evidence first = evidence( + "evidence-1", + "Alice is the recorded name"); + List evidence = new ArrayList<>(); + evidence.add(first); + + MemoryFactVersion version = new MemoryFactVersion( + "version-1", + fact("Alice"), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + TimeInterval.unboundedFrom( + time("2024-03-01T00:00:00Z")), + evidence); + + evidence.add(evidence( + "evidence-2", + "A later independent record")); + + Assertions.assertEquals( + Collections.singletonList(first), + version.getEvidence()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> version.getEvidence().clear()); + } + + @Test + public void testStatusDefaultsToActiveAndAffectsValueSemantics() { + MemoryFact fact = fact("Alice"); + TimeInterval validTime = TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")); + TimeInterval transactionTime = TimeInterval.unboundedFrom( + time("2024-03-01T00:00:00Z")); + List evidence = Collections.singletonList( + evidence("evidence-1", "Alice is the recorded name")); + MemoryFactVersion active = new MemoryFactVersion( + "version-1", + fact, + validTime, + transactionTime, + evidence); + MemoryFactVersion explicitActive = new MemoryFactVersion( + "version-1", + fact, + MemoryFactVersionStatus.ACTIVE, + validTime, + transactionTime, + evidence); + MemoryFactVersion retracted = new MemoryFactVersion( + "version-1", + fact, + MemoryFactVersionStatus.RETRACTED, + validTime, + transactionTime, + evidence); + + Assertions.assertEquals( + MemoryFactVersionStatus.ACTIVE, + active.getStatus()); + Assertions.assertEquals(active, explicitActive); + Assertions.assertNotEquals(active, retracted); + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryFactVersion( + "version-1", + fact, + null, + validTime, + transactionTime, + evidence)); + } + + @Test + public void testRejectInvalidVersion() { + MemoryFact fact = fact("Alice"); + TimeInterval validTime = TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")); + TimeInterval transactionTime = TimeInterval.unboundedFrom( + time("2024-03-01T00:00:00Z")); + List evidence = Collections.singletonList( + evidence("evidence-1", "Alice is the recorded name")); + + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryFactVersion( + null, + fact, + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new MemoryFactVersion( + " ", + fact, + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryFactVersion( + "version-1", + null, + validTime, + transactionTime, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryFactVersion( + "version-1", + fact, + null, + transactionTime, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryFactVersion( + "version-1", + fact, + validTime, + null, + evidence)); + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryFactVersion( + "version-1", + fact, + validTime, + transactionTime, + null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new MemoryFactVersion( + "version-1", + fact, + validTime, + transactionTime, + Collections.emptyList())); + + List evidenceWithNull = new ArrayList<>(); + evidenceWithNull.add(null); + + Assertions.assertThrows( + NullPointerException.class, + () -> new MemoryFactVersion( + "version-1", + fact, + validTime, + transactionTime, + evidenceWithNull)); + } + + private static MemoryFact fact(String literalValue) { + return MemoryFact.attribute( + "fact-name-alice", + new MemoryEntity("person:alice", "person"), + "name", + literalValue); + } + + private static Evidence evidence(String id, String content) { + return new Evidence( + id, + new Source("source-1", "customer-database"), + content); + } + + private static TimeInterval interval(String start, String end) { + return new TimeInterval(time(start), time(end)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/ProvenanceModelTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/ProvenanceModelTest.java new file mode 100644 index 000000000..a233690e9 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/ProvenanceModelTest.java @@ -0,0 +1,89 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class ProvenanceModelTest { + + @Test + public void testSourceValueSemantics() { + Source source = new Source("source-1", "customer-database"); + Source same = new Source("source-1", "customer-database"); + + Assertions.assertEquals("source-1", source.getId()); + Assertions.assertEquals("customer-database", source.getName()); + Assertions.assertEquals(source, same); + Assertions.assertEquals(source.hashCode(), same.hashCode()); + Assertions.assertNotEquals( + source, + new Source("source-1", "archive-database")); + } + + @Test + public void testEvidenceValueSemantics() { + Source source = new Source("source-1", "customer-database"); + Evidence evidence = new Evidence( + "evidence-1", + source, + "Alice works at Acme"); + Evidence same = new Evidence( + "evidence-1", + new Source("source-1", "customer-database"), + "Alice works at Acme"); + + Assertions.assertEquals("evidence-1", evidence.getId()); + Assertions.assertEquals(source, evidence.getSource()); + Assertions.assertEquals( + "Alice works at Acme", + evidence.getContent()); + Assertions.assertEquals(evidence, same); + Assertions.assertEquals(evidence.hashCode(), same.hashCode()); + Assertions.assertNotEquals( + evidence, + new Evidence("evidence-1", source, "Alice left Acme")); + } + + @Test + public void testRejectInvalidProvenance() { + Source source = new Source("source-1", "customer-database"); + + Assertions.assertThrows( + NullPointerException.class, + () -> new Source(null, "customer-database")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new Source(" ", "customer-database")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new Source("source-1", " ")); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new Evidence(" ", source, "Alice works at Acme")); + Assertions.assertThrows( + NullPointerException.class, + () -> new Evidence("evidence-1", null, "Alice works at Acme")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new Evidence("evidence-1", source, " ")); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/TimeIntervalTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/TimeIntervalTest.java new file mode 100644 index 000000000..e0efd94de --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/TimeIntervalTest.java @@ -0,0 +1,139 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.time.Instant; +import java.util.Arrays; +import java.util.Collections; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class TimeIntervalTest { + + @Test + public void testContainsUsesHalfOpenBounds() { + TimeInterval interval = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + + Assertions.assertTrue(interval.contains(time("2024-01-01T00:00:00Z"))); + Assertions.assertTrue(interval.contains(time("2024-06-01T00:00:00Z"))); + Assertions.assertFalse(interval.contains(time("2025-01-01T00:00:00Z"))); + Assertions.assertFalse(interval.contains(time("2023-12-31T23:59:59Z"))); + } + + @Test + public void testUnboundedInterval() { + TimeInterval interval = + TimeInterval.unboundedFrom(time("2024-01-01T00:00:00Z")); + + Assertions.assertFalse(interval.getEnd().isPresent()); + Assertions.assertTrue(interval.contains(time("2099-01-01T00:00:00Z"))); + Assertions.assertFalse(interval.contains(time("2023-01-01T00:00:00Z"))); + } + + @Test + public void testIntersectionAndOverlap() { + TimeInterval left = interval( + "2024-01-01T00:00:00Z", + "2024-10-01T00:00:00Z"); + TimeInterval right = interval( + "2024-05-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + TimeInterval touching = interval( + "2024-10-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + + Assertions.assertTrue(left.overlaps(right)); + Assertions.assertEquals( + interval("2024-05-01T00:00:00Z", "2024-10-01T00:00:00Z"), + left.intersection(right).get()); + + Assertions.assertFalse(left.overlaps(touching)); + Assertions.assertFalse(left.intersection(touching).isPresent()); + } + + @Test + public void testSubtract() { + TimeInterval whole = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + TimeInterval middle = interval( + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z"); + + Assertions.assertEquals( + Arrays.asList( + interval("2024-01-01T00:00:00Z", "2024-04-01T00:00:00Z"), + interval("2024-09-01T00:00:00Z", "2025-01-01T00:00:00Z")), + whole.subtract(middle)); + + Assertions.assertEquals( + Collections.singletonList(whole), + whole.subtract(interval( + "2025-01-01T00:00:00Z", + "2026-01-01T00:00:00Z"))); + + Assertions.assertTrue(whole.subtract(whole).isEmpty()); + } + + @Test + public void testRejectInvalidInterval() { + Instant start = time("2024-01-01T00:00:00Z"); + + Assertions.assertThrows( + NullPointerException.class, + () -> new TimeInterval(null, start)); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new TimeInterval(start, start)); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new TimeInterval( + time("2025-01-01T00:00:00Z"), + time("2024-01-01T00:00:00Z"))); + } + + @Test + public void testSubtractFromUnboundedInterval() { + TimeInterval whole = + TimeInterval.unboundedFrom(time("2024-01-01T00:00:00Z")); + TimeInterval removed = interval( + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z"); + + Assertions.assertEquals( + Arrays.asList( + interval("2024-01-01T00:00:00Z", "2024-04-01T00:00:00Z"), + TimeInterval.unboundedFrom(time("2024-09-01T00:00:00Z"))), + whole.subtract(removed)); + } + + private static TimeInterval interval(String start, String end) { + return new TimeInterval(time(start), time(end)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/VersionRelationTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/VersionRelationTest.java new file mode 100644 index 000000000..0a0f9c9f4 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/model/VersionRelationTest.java @@ -0,0 +1,143 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.model; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class VersionRelationTest { + + @Test + public void testDirectionalRelationValueSemantics() { + VersionRelation relation = new VersionRelation( + VersionRelationType.SUPERSEDES, + "version-2", + "version-1"); + VersionRelation same = new VersionRelation( + VersionRelationType.SUPERSEDES, + "version-2", + "version-1"); + + Assertions.assertEquals( + VersionRelationType.SUPERSEDES, + relation.getType()); + Assertions.assertEquals( + "version-2", + relation.getFromVersionId()); + Assertions.assertEquals( + "version-1", + relation.getToVersionId()); + Assertions.assertEquals(relation, same); + Assertions.assertEquals( + relation.hashCode(), + same.hashCode()); + } + + @Test + public void testConflictUsesCanonicalDirection() { + VersionRelation relation = new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + "version-b", + "version-a"); + VersionRelation canonical = new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + "version-a", + "version-b"); + + Assertions.assertEquals( + "version-a", + relation.getFromVersionId()); + Assertions.assertEquals( + "version-b", + relation.getToVersionId()); + Assertions.assertEquals(canonical, relation); + } + + @Test + public void testDeterministicOrder() { + VersionRelation supersedes = new VersionRelation( + VersionRelationType.SUPERSEDES, + "version-2", + "version-1"); + VersionRelation duplicate = new VersionRelation( + VersionRelationType.DUPLICATE_OF, + "version-3", + "version-1"); + VersionRelation conflict = new VersionRelation( + VersionRelationType.CONFLICTS_WITH, + "version-2", + "version-3"); + List relations = + new ArrayList<>(Arrays.asList( + conflict, + duplicate, + supersedes)); + + Collections.sort(relations); + + Assertions.assertEquals( + Arrays.asList(supersedes, duplicate, conflict), + relations); + } + + @Test + public void testRejectInvalidRelation() { + Assertions.assertThrows( + NullPointerException.class, + () -> new VersionRelation( + null, + "version-2", + "version-1")); + Assertions.assertThrows( + NullPointerException.class, + () -> new VersionRelation( + VersionRelationType.SUPERSEDES, + null, + "version-1")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new VersionRelation( + VersionRelationType.SUPERSEDES, + " ", + "version-1")); + Assertions.assertThrows( + NullPointerException.class, + () -> new VersionRelation( + VersionRelationType.SUPERSEDES, + "version-2", + null)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new VersionRelation( + VersionRelationType.SUPERSEDES, + "version-2", + " ")); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> new VersionRelation( + VersionRelationType.SUPERSEDES, + "version-1", + "version-1")); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracleTemporalStateTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracleTemporalStateTest.java new file mode 100644 index 000000000..0840f726c --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracleTemporalStateTest.java @@ -0,0 +1,484 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.oracle; + +import java.time.Instant; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.apache.geaflow.ai.temporal.query.BitemporalQuery; +import org.apache.geaflow.ai.temporal.semantics.EventNormalizer; +import org.apache.geaflow.ai.temporal.semantics.NormalizedMemoryEvent; +import org.apache.geaflow.ai.temporal.semantics.TemporalState; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class FullReplayOracleTemporalStateTest { + + private static final FactKey KEY = new FactKey( + "person:alice", + "city", + "profile"); + + private final EventNormalizer normalizer = new EventNormalizer(); + private final FullReplayOracle oracle = new FullReplayOracle(); + private final BitemporalQuery query = new BitemporalQuery(); + + @Test + public void testPartialRetractionCreatesTombstoneAndResidues() { + NormalizedMemoryEvent add = add( + "event-add", + KEY, + "Beijing", + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent retract = retract( + "event-retract", + KEY, + interval( + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z"), + "2024-06-01T00:00:00Z"); + + TemporalState state = oracle.replayNormalized( + Arrays.asList(retract, add)); + + Assertions.assertEquals(4, state.getVersions().size()); + Assertions.assertEquals( + state.getVersions(), + state.getVersionsByFactKey().get(KEY)); + + MemoryFactVersion old = version( + state, + "event-add:version:0"); + Assertions.assertEquals( + MemoryFactVersionStatus.ACTIVE, + old.getStatus()); + Assertions.assertEquals( + interval( + "2024-03-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + old.getTransactionTime()); + + MemoryFactVersion tombstone = version( + state, + "event-retract:version:0"); + Assertions.assertEquals( + MemoryFactVersionStatus.RETRACTED, + tombstone.getStatus()); + Assertions.assertEquals( + retract.getValidTime(), + tombstone.getValidTime()); + Assertions.assertEquals( + TimeInterval.unboundedFrom(retract.getRecordedAt()), + tombstone.getTransactionTime()); + + MemoryFactVersion left = version( + state, + "event-retract:version:1"); + MemoryFactVersion right = version( + state, + "event-retract:version:2"); + Assertions.assertEquals( + MemoryFactVersionStatus.ACTIVE, + left.getStatus()); + Assertions.assertEquals( + interval( + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"), + left.getValidTime()); + Assertions.assertEquals( + MemoryFactVersionStatus.ACTIVE, + right.getStatus()); + Assertions.assertEquals( + TimeInterval.unboundedFrom( + time("2024-09-01T00:00:00Z")), + right.getValidTime()); + + Assertions.assertEquals( + Arrays.asList( + relation( + VersionRelationType.SUPERSEDES, + tombstone.getId(), + old.getId()), + relation( + VersionRelationType.SUPERSEDES, + left.getId(), + old.getId()), + relation( + VersionRelationType.SUPERSEDES, + right.getId(), + old.getId())), + state.getRelations()); + Assertions.assertTrue(query.query( + state.getVersions(), + time("2024-05-01T00:00:00Z"), + time("2024-07-01T00:00:00Z")).isEmpty()); + Assertions.assertEquals(1, query.query( + state.getVersions(), + time("2024-02-01T00:00:00Z"), + time("2024-07-01T00:00:00Z")).size()); + } + + @Test + public void testExplicitCorrectionSupersedesSameValue() { + TimeInterval validTime = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + NormalizedMemoryEvent add = add( + "event-add", + KEY, + "Beijing", + validTime, + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent correction = correct( + "event-correct", + KEY, + "Beijing", + validTime, + "2024-06-01T00:00:00Z"); + + TemporalState state = oracle.replayNormalized( + Arrays.asList(correction, add)); + + Assertions.assertEquals(2, state.getVersions().size()); + Assertions.assertEquals( + Collections.singletonList(relation( + VersionRelationType.SUPERSEDES, + "event-correct:version:0", + "event-add:version:0")), + state.getRelations()); + } + + @Test + public void testSameValueOverlapCreatesDuplicateVersion() { + TimeInterval validTime = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + NormalizedMemoryEvent first = add( + "event-first", + KEY, + "Beijing", + validTime, + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent duplicate = add( + "event-duplicate", + KEY, + "Beijing", + validTime, + "2024-04-01T00:00:00Z"); + + TemporalState state = oracle.replayNormalized( + Arrays.asList(duplicate, first)); + + MemoryFactVersion old = version( + state, + "event-first:version:0"); + MemoryFactVersion current = version( + state, + "event-duplicate:version:0"); + Assertions.assertEquals( + time("2024-04-01T00:00:00Z"), + old.getTransactionTime().getEnd().get()); + Assertions.assertFalse( + current.getTransactionTime().getEnd().isPresent()); + Assertions.assertEquals( + Arrays.asList( + "evidence-event-duplicate", + "evidence-event-first"), + Arrays.asList( + current.getEvidence().get(0).getId(), + current.getEvidence().get(1).getId())); + Assertions.assertEquals( + Collections.singletonList(relation( + VersionRelationType.DUPLICATE_OF, + current.getId(), + old.getId())), + state.getRelations()); + } + + @Test + public void testDifferentValueOverlapConflictsButBoundaryDoesNot() { + NormalizedMemoryEvent first = add( + "event-first", + KEY, + "Beijing", + interval( + "2024-01-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent conflict = add( + "event-conflict", + KEY, + "Shanghai", + interval( + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z"), + "2024-04-01T00:00:00Z"); + NormalizedMemoryEvent boundary = add( + "event-boundary", + KEY, + "Rome", + interval( + "2024-09-01T00:00:00Z", + "2024-12-01T00:00:00Z"), + "2024-05-01T00:00:00Z"); + + TemporalState shuffled = oracle.replayNormalized( + Arrays.asList(boundary, conflict, first)); + TemporalState ordered = oracle.replayNormalized( + Arrays.asList(first, conflict, boundary)); + + Assertions.assertEquals(ordered, shuffled); + Assertions.assertEquals(3, shuffled.getVersions().size()); + for (MemoryFactVersion version : shuffled.getVersions()) { + Assertions.assertEquals( + MemoryFactVersionStatus.ACTIVE, + version.getStatus()); + Assertions.assertFalse( + version.getTransactionTime().getEnd().isPresent()); + } + Assertions.assertEquals( + Collections.singletonList(relation( + VersionRelationType.CONFLICTS_WITH, + "event-conflict:version:0", + "event-first:version:0")), + shuffled.getRelations()); + } + + @Test + public void testPartialChangesAddConflictsForResidualVersions() { + NormalizedMemoryEvent first = add( + "event-first", + KEY, + "Beijing", + interval( + "2024-01-01T00:00:00Z", + "2024-10-01T00:00:00Z"), + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent conflict = add( + "event-conflict", + KEY, + "Shanghai", + interval( + "2024-05-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + "2024-04-01T00:00:00Z"); + NormalizedMemoryEvent retraction = retract( + "event-retract", + KEY, + interval( + "2024-01-01T00:00:00Z", + "2024-05-01T00:00:00Z"), + "2024-06-01T00:00:00Z"); + NormalizedMemoryEvent correction = correct( + "event-correct", + KEY, + "Rome", + retraction.getValidTime(), + "2024-06-01T00:00:00Z"); + + TemporalState state = oracle.replayNormalized( + Arrays.asList(retraction, conflict, first)); + + Assertions.assertTrue(state.getRelations().contains(relation( + VersionRelationType.CONFLICTS_WITH, + "event-conflict:version:0", + "event-first:version:0"))); + Assertions.assertTrue(state.getRelations().contains(relation( + VersionRelationType.CONFLICTS_WITH, + "event-conflict:version:0", + "event-retract:version:1"))); + + TemporalState correctedState = oracle.replayNormalized( + Arrays.asList(correction, conflict, first)); + Assertions.assertTrue(correctedState.getRelations().contains( + relation( + VersionRelationType.CONFLICTS_WITH, + "event-conflict:version:0", + "event-correct:version:1"))); + } + + @Test + public void testLedgerNoopAndEventIdReuse() { + NormalizedMemoryEvent original = add( + "event-1", + KEY, + "Beijing", + interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent reused = add( + "event-1", + KEY, + "Shanghai", + interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"), + "2024-03-01T00:00:00Z"); + + Assertions.assertEquals( + oracle.replayNormalized( + Collections.singletonList(original)), + oracle.replayNormalized( + Arrays.asList(original, original))); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> oracle.replayNormalized( + Arrays.asList(original, reused))); + } + + @Test + public void testSameRecordedAtDoesNotLeaveZeroLengthVersion() { + TimeInterval validTime = interval( + "2024-01-01T00:00:00Z", + "2025-01-01T00:00:00Z"); + NormalizedMemoryEvent add = add( + "event-a-add", + KEY, + "Beijing", + validTime, + "2024-03-01T00:00:00Z"); + NormalizedMemoryEvent correction = correct( + "event-b-correct", + KEY, + "Shanghai", + validTime, + "2024-03-01T00:00:00Z"); + + TemporalState state = oracle.replayNormalized( + Arrays.asList(correction, add)); + + Assertions.assertEquals(1, state.getVersions().size()); + Assertions.assertEquals( + "event-b-correct:version:0", + state.getVersions().get(0).getId()); + Assertions.assertTrue(state.getRelations().isEmpty()); + } + + private NormalizedMemoryEvent add( + String eventId, + FactKey key, + String value, + TimeInterval validTime, + String recordedAt) { + return normalizer.normalize( + MemoryEvent.add( + eventId, + fact(key, value), + validTime, + time(recordedAt), + evidence(eventId)), + key); + } + + private NormalizedMemoryEvent correct( + String eventId, + FactKey key, + String value, + TimeInterval validTime, + String recordedAt) { + return normalizer.normalize( + MemoryEvent.correct( + eventId, + fact(key, value), + validTime, + time(recordedAt), + evidence(eventId)), + key); + } + + private NormalizedMemoryEvent retract( + String eventId, + FactKey key, + TimeInterval validTime, + String recordedAt) { + return normalizer.normalize( + MemoryEvent.retract( + eventId, + factId(key), + validTime, + time(recordedAt), + evidence(eventId)), + key); + } + + private static MemoryFact fact( + FactKey key, + String value) { + return MemoryFact.attribute( + factId(key), + new MemoryEntity(key.getSubjectId(), "person"), + key.getPredicate(), + value); + } + + private static String factId(FactKey key) { + return "fact-city-" + key.getScope(); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "registry"), + "Evidence for " + eventId)); + } + + private static VersionRelation relation( + VersionRelationType type, + String from, + String to) { + return new VersionRelation(type, from, to); + } + + private static MemoryFactVersion version( + TemporalState state, + String versionId) { + for (MemoryFactVersion version : state.getVersions()) { + if (version.getId().equals(versionId)) { + return version; + } + } + throw new AssertionError("Missing version: " + versionId); + } + + private static TimeInterval interval( + String start, + String end) { + return new TimeInterval(time(start), time(end)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracleTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracleTest.java new file mode 100644 index 000000000..988f8b3eb --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/oracle/FullReplayOracleTest.java @@ -0,0 +1,627 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.oracle; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class FullReplayOracleTest { + + private final FullReplayOracle oracle = new FullReplayOracle(); + + @Test + public void testReplayAddEvent() { + MemoryEvent event = addEvent( + "event-1", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + + List versions = + oracle.replay(Collections.singletonList(event)); + + Assertions.assertEquals(1, versions.size()); + + MemoryFactVersion version = versions.get(0); + Assertions.assertEquals( + "event-1:version:0", + version.getId()); + Assertions.assertEquals( + event.getFact().get(), + version.getFact()); + Assertions.assertEquals( + event.getValidTime(), + version.getValidTime()); + Assertions.assertEquals( + TimeInterval.unboundedFrom( + event.getTransactionTime()), + version.getTransactionTime()); + Assertions.assertEquals( + event.getEvidence(), + version.getEvidence()); + } + + @Test + public void testReplayUsesDeterministicOrder() { + MemoryEvent eventB = addEvent( + "event-b", + "fact-name-bob", + "person:bob", + "Bob", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent eventA = addEvent( + "event-a", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent eventC = addEvent( + "event-c", + "fact-name-carol", + "person:carol", + "Carol", + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"); + + List shuffled = oracle.replay( + Arrays.asList(eventC, eventB, eventA)); + List ordered = oracle.replay( + Arrays.asList(eventA, eventB, eventC)); + + Assertions.assertEquals(ordered, shuffled); + Assertions.assertEquals( + "event-a:version:0", + shuffled.get(0).getId()); + Assertions.assertEquals( + "event-b:version:0", + shuffled.get(1).getId()); + Assertions.assertEquals( + "event-c:version:0", + shuffled.get(2).getId()); + } + + @Test + public void testDuplicateEventIsIdempotent() { + MemoryEvent event = addEvent( + "event-1", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent duplicate = addEvent( + "event-1", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + + List once = + oracle.replay(Collections.singletonList(event)); + List repeated = + oracle.replay(Arrays.asList(event, duplicate, event)); + + Assertions.assertEquals(once, repeated); + Assertions.assertEquals(1, repeated.size()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> repeated.clear()); + } + + @Test + public void testRejectConflictingEventId() { + MemoryEvent original = addEvent( + "event-1", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent conflict = addEvent( + "event-1", + "fact-name-alice", + "person:alice", + "Alice Smith", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> oracle.replay( + Arrays.asList(original, conflict))); + } + + @Test + public void testRetractEventSplitsValidTimeAndClosesOldVersion() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent retract = retractEvent( + "event-retract", + "fact-name-alice", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + + List versions = + oracle.replay(Arrays.asList(retract, add)); + + Assertions.assertEquals(3, versions.size()); + assertVersion( + versions.get(0), + "event-add:version:0", + add.getFact().get(), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + interval( + "2024-03-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + add.getEvidence()); + assertVersion( + versions.get(1), + "event-retract:version:1", + add.getFact().get(), + interval( + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"), + TimeInterval.unboundedFrom( + time("2024-06-01T00:00:00Z")), + add.getEvidence()); + assertVersion( + versions.get(2), + "event-retract:version:2", + add.getFact().get(), + TimeInterval.unboundedFrom( + time("2024-09-01T00:00:00Z")), + TimeInterval.unboundedFrom( + time("2024-06-01T00:00:00Z")), + add.getEvidence()); + } + + @Test + public void testRetractWholeCurrentInterval() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent retract = MemoryEvent.retract( + "event-retract", + "fact-name-alice", + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time("2024-06-01T00:00:00Z"), + evidence("event-retract")); + + List versions = + oracle.replay(Arrays.asList(add, retract)); + + Assertions.assertEquals(1, versions.size()); + Assertions.assertEquals( + interval( + "2024-03-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + versions.get(0).getTransactionTime()); + } + + @Test + public void testCorrectEventSplitsValidTimeAndClosesOldVersion() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent correction = correctEvent( + "event-correct", + "fact-name-alice", + "person:alice", + "Alice Smith", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + + List versions = + oracle.replay(Arrays.asList(correction, add)); + + Assertions.assertEquals(4, versions.size()); + + assertVersion( + versions.get(0), + "event-add:version:0", + add.getFact().get(), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + interval( + "2024-03-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + add.getEvidence()); + assertVersion( + versions.get(1), + "event-correct:version:1", + add.getFact().get(), + interval( + "2024-01-01T00:00:00Z", + "2024-04-01T00:00:00Z"), + TimeInterval.unboundedFrom( + time("2024-06-01T00:00:00Z")), + add.getEvidence()); + assertVersion( + versions.get(2), + "event-correct:version:0", + correction.getFact().get(), + interval( + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z"), + TimeInterval.unboundedFrom( + time("2024-06-01T00:00:00Z")), + correction.getEvidence()); + assertVersion( + versions.get(3), + "event-correct:version:2", + add.getFact().get(), + TimeInterval.unboundedFrom( + time("2024-09-01T00:00:00Z")), + TimeInterval.unboundedFrom( + time("2024-06-01T00:00:00Z")), + add.getEvidence()); + } + + @Test + public void testCorrectAcrossCurrentFragments() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent firstCorrection = correctEvent( + "event-correct-1", + "fact-name-alice", + "person:alice", + "Alice Smith", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + MemoryEvent secondCorrection = correctEvent( + "event-correct-2", + "fact-name-alice", + "person:alice", + "Alice Jones", + "2024-02-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-08-01T00:00:00Z"); + + List versions = oracle.replay( + Arrays.asList( + secondCorrection, + add, + firstCorrection)); + + List current = new ArrayList<>(); + for (MemoryFactVersion version : versions) { + if (!version.getTransactionTime() + .getEnd().isPresent()) { + current.add(version); + } + } + + Assertions.assertEquals(3, current.size()); + Assertions.assertEquals( + "event-correct-2:version:1", + current.get(0).getId()); + Assertions.assertEquals( + add.getFact().get(), + current.get(0).getFact()); + Assertions.assertEquals( + interval( + "2024-01-01T00:00:00Z", + "2024-02-01T00:00:00Z"), + current.get(0).getValidTime()); + + Assertions.assertEquals( + "event-correct-2:version:0", + current.get(1).getId()); + Assertions.assertEquals( + secondCorrection.getFact().get(), + current.get(1).getFact()); + Assertions.assertEquals( + interval( + "2024-02-01T00:00:00Z", + "2024-10-01T00:00:00Z"), + current.get(1).getValidTime()); + + Assertions.assertEquals( + "event-correct-2:version:2", + current.get(2).getId()); + Assertions.assertEquals( + add.getFact().get(), + current.get(2).getFact()); + Assertions.assertEquals( + TimeInterval.unboundedFrom( + time("2024-10-01T00:00:00Z")), + current.get(2).getValidTime()); + } + + @Test + public void testRetractAcrossCurrentFragments() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent correction = correctEvent( + "event-correct", + "fact-name-alice", + "person:alice", + "Alice Smith", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + MemoryEvent retract = retractEvent( + "event-retract", + "fact-name-alice", + "2024-02-01T00:00:00Z", + "2024-10-01T00:00:00Z", + "2024-08-01T00:00:00Z"); + + List versions = oracle.replay( + Arrays.asList(retract, correction, add)); + + List current = new ArrayList<>(); + for (MemoryFactVersion version : versions) { + if (!version.getTransactionTime() + .getEnd().isPresent()) { + current.add(version); + } + } + + Assertions.assertEquals(6, versions.size()); + Assertions.assertEquals(2, current.size()); + Assertions.assertEquals( + "event-retract:version:1", + current.get(0).getId()); + Assertions.assertEquals( + add.getFact().get(), + current.get(0).getFact()); + Assertions.assertEquals( + interval( + "2024-01-01T00:00:00Z", + "2024-02-01T00:00:00Z"), + current.get(0).getValidTime()); + Assertions.assertEquals( + "event-retract:version:2", + current.get(1).getId()); + Assertions.assertEquals( + add.getFact().get(), + current.get(1).getFact()); + Assertions.assertEquals( + TimeInterval.unboundedFrom( + time("2024-10-01T00:00:00Z")), + current.get(1).getValidTime()); + } + + @Test + public void testRejectCorrectionWithoutFullCoverage() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent correction = correctEvent( + "event-correct", + "fact-name-alice", + "person:alice", + "Alice Smith", + "2023-12-01T00:00:00Z", + "2024-02-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> oracle.replay( + Arrays.asList(add, correction))); + } + + @Test + public void testRejectRetractionWithoutFullCoverage() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent retract = retractEvent( + "event-retract", + "fact-name-alice", + "2023-12-01T00:00:00Z", + "2024-02-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> oracle.replay( + Arrays.asList(add, retract))); + } + + @Test + public void testRejectOverlappingAddForSameFact() { + MemoryEvent first = addEvent( + "event-1", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent second = addEvent( + "event-2", + "fact-name-alice", + "person:alice", + "Alice Smith", + "2024-02-01T00:00:00Z", + "2024-04-01T00:00:00Z"); + + Assertions.assertThrows( + IllegalArgumentException.class, + () -> oracle.replay(Arrays.asList(first, second))); + } + + @Test + public void testEmptyAndInvalidInput() { + Assertions.assertTrue( + oracle.replay(Collections.emptyList()).isEmpty()); + Assertions.assertThrows( + NullPointerException.class, + () -> oracle.replay(null)); + + List eventsWithNull = new ArrayList<>(); + eventsWithNull.add(null); + + Assertions.assertThrows( + NullPointerException.class, + () -> oracle.replay(eventsWithNull)); + } + + private static MemoryEvent addEvent( + String eventId, + String factId, + String subjectId, + String literalValue, + String validStart, + String transactionTime) { + MemoryFact fact = MemoryFact.attribute( + factId, + new MemoryEntity(subjectId, "person"), + "name", + literalValue); + + return MemoryEvent.add( + eventId, + fact, + TimeInterval.unboundedFrom(time(validStart)), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent correctEvent( + String eventId, + String factId, + String subjectId, + String literalValue, + String validStart, + String validEnd, + String transactionTime) { + MemoryFact fact = MemoryFact.attribute( + factId, + new MemoryEntity(subjectId, "person"), + "name", + literalValue); + + return MemoryEvent.correct( + eventId, + fact, + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent retractEvent( + String eventId, + String factId, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.retract( + eventId, + factId, + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static void assertVersion( + MemoryFactVersion actual, + String expectedId, + MemoryFact expectedFact, + TimeInterval expectedValidTime, + TimeInterval expectedTransactionTime, + List expectedEvidence) { + Assertions.assertEquals(expectedId, actual.getId()); + Assertions.assertEquals( + expectedFact, + actual.getFact()); + Assertions.assertEquals( + expectedValidTime, + actual.getValidTime()); + Assertions.assertEquals( + expectedTransactionTime, + actual.getTransactionTime()); + Assertions.assertEquals( + expectedEvidence, + actual.getEvidence()); + } + + private static TimeInterval interval( + String start, + String end) { + return new TimeInterval(time(start), time(end)); + } + + private static List evidence(String eventId) { + return Collections.singletonList(new Evidence( + "evidence-" + eventId, + new Source("source-1", "customer-database"), + "Evidence for " + eventId)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/query/BitemporalQueryTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/query/BitemporalQueryTest.java new file mode 100644 index 000000000..180b18d7c --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/query/BitemporalQueryTest.java @@ -0,0 +1,303 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.query; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.oracle.FullReplayOracle; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class BitemporalQueryTest { + + private final BitemporalQuery query = new BitemporalQuery(); + private final FullReplayOracle oracle = new FullReplayOracle(); + + @Test + public void testQueryBeforeAndAfterCorrection() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent correction = correctEvent( + "event-correct", + "fact-name-alice", + "person:alice", + "Alice Smith", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + List versions = + oracle.replay(Arrays.asList(correction, add)); + + List beforeCorrection = query.query( + versions, + time("2024-05-01T00:00:00Z"), + time("2024-05-01T00:00:00Z")); + List afterCorrection = query.query( + versions, + time("2024-05-01T00:00:00Z"), + time("2024-07-01T00:00:00Z")); + + Assertions.assertEquals(1, beforeCorrection.size()); + Assertions.assertEquals( + add.getFact().get(), + beforeCorrection.get(0).getFact()); + Assertions.assertEquals(1, afterCorrection.size()); + Assertions.assertEquals( + correction.getFact().get(), + afterCorrection.get(0).getFact()); + } + + @Test + public void testQueryBeforeAndAfterRetraction() { + MemoryEvent add = addEvent( + "event-add", + "fact-name-alice", + "person:alice", + "Alice", + "2024-01-01T00:00:00Z", + "2024-03-01T00:00:00Z"); + MemoryEvent retract = retractEvent( + "event-retract", + "fact-name-alice", + "2024-04-01T00:00:00Z", + "2024-09-01T00:00:00Z", + "2024-06-01T00:00:00Z"); + List versions = + oracle.replay(Arrays.asList(retract, add)); + + List beforeRetraction = query.query( + versions, + time("2024-05-01T00:00:00Z"), + time("2024-05-01T00:00:00Z")); + List afterRetraction = query.query( + versions, + time("2024-05-01T00:00:00Z"), + time("2024-07-01T00:00:00Z")); + List outsideRetraction = query.query( + versions, + time("2024-10-01T00:00:00Z"), + time("2024-07-01T00:00:00Z")); + + Assertions.assertEquals(1, beforeRetraction.size()); + Assertions.assertEquals( + add.getFact().get(), + beforeRetraction.get(0).getFact()); + Assertions.assertTrue(afterRetraction.isEmpty()); + Assertions.assertEquals(1, outsideRetraction.size()); + Assertions.assertEquals( + add.getFact().get(), + outsideRetraction.get(0).getFact()); + } + + @Test + public void testHalfOpenBoundaries() { + MemoryFactVersion version = new MemoryFactVersion( + "version-1", + fact("fact-name-alice", "person:alice", "Alice"), + interval( + "2024-01-01T00:00:00Z", + "2024-06-01T00:00:00Z"), + interval( + "2024-03-01T00:00:00Z", + "2024-09-01T00:00:00Z"), + evidence("version-1")); + List versions = + Collections.singletonList(version); + + Assertions.assertEquals( + 1, + query.query( + versions, + time("2024-01-01T00:00:00Z"), + time("2024-03-01T00:00:00Z")).size()); + Assertions.assertTrue( + query.query( + versions, + time("2024-06-01T00:00:00Z"), + time("2024-03-01T00:00:00Z")).isEmpty()); + Assertions.assertTrue( + query.query( + versions, + time("2024-05-01T00:00:00Z"), + time("2024-09-01T00:00:00Z")).isEmpty()); + } + + @Test + public void testResultIsDeterministicAndImmutable() { + MemoryFactVersion bob = currentVersion( + "version-bob", + fact("fact-name-bob", "person:bob", "Bob")); + MemoryFactVersion alice = currentVersion( + "version-alice", + fact("fact-name-alice", "person:alice", "Alice")); + + List result = query.query( + Arrays.asList(bob, alice), + time("2024-05-01T00:00:00Z"), + time("2024-05-01T00:00:00Z")); + + Assertions.assertEquals(2, result.size()); + Assertions.assertEquals( + "fact-name-alice", + result.get(0).getFact().getId()); + Assertions.assertEquals( + "fact-name-bob", + result.get(1).getFact().getId()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> result.clear()); + } + + @Test + public void testEmptyAndInvalidInput() { + Instant queryTime = time("2024-05-01T00:00:00Z"); + + Assertions.assertTrue( + query.query( + Collections.emptyList(), + queryTime, + queryTime).isEmpty()); + Assertions.assertThrows( + NullPointerException.class, + () -> query.query(null, queryTime, queryTime)); + Assertions.assertThrows( + NullPointerException.class, + () -> query.query( + Collections.emptyList(), + null, + queryTime)); + Assertions.assertThrows( + NullPointerException.class, + () -> query.query( + Collections.emptyList(), + queryTime, + null)); + + List versionsWithNull = new ArrayList<>(); + versionsWithNull.add(null); + Assertions.assertThrows( + NullPointerException.class, + () -> query.query( + versionsWithNull, + queryTime, + queryTime)); + } + + private static MemoryEvent addEvent( + String eventId, + String factId, + String subjectId, + String literalValue, + String validStart, + String transactionTime) { + return MemoryEvent.add( + eventId, + fact(factId, subjectId, literalValue), + TimeInterval.unboundedFrom(time(validStart)), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent correctEvent( + String eventId, + String factId, + String subjectId, + String literalValue, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.correct( + eventId, + fact(factId, subjectId, literalValue), + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryEvent retractEvent( + String eventId, + String factId, + String validStart, + String validEnd, + String transactionTime) { + return MemoryEvent.retract( + eventId, + factId, + interval(validStart, validEnd), + time(transactionTime), + evidence(eventId)); + } + + private static MemoryFactVersion currentVersion( + String versionId, + MemoryFact fact) { + return new MemoryFactVersion( + versionId, + fact, + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + TimeInterval.unboundedFrom( + time("2024-03-01T00:00:00Z")), + evidence(versionId)); + } + + private static MemoryFact fact( + String factId, + String subjectId, + String literalValue) { + return MemoryFact.attribute( + factId, + new MemoryEntity(subjectId, "person"), + "name", + literalValue); + } + + private static TimeInterval interval( + String start, + String end) { + return new TimeInterval(time(start), time(end)); + } + + private static List evidence(String id) { + return Collections.singletonList(new Evidence( + "evidence-" + id, + new Source("source-1", "customer-database"), + "Evidence for " + id)); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/CanonicalSnapshotTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/CanonicalSnapshotTest.java new file mode 100644 index 000000000..717c10cce --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/CanonicalSnapshotTest.java @@ -0,0 +1,348 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +/** + * Preserves normalized payloads and explicit version-generation references. + * Events use recordedAt/ID order; state keeps FactKey/system-time/valid-time/ID order. + * Version evidence is sorted without dropping repeated entries. + */ +public class CanonicalSnapshotTest { + + private static final FactKey KEY = key("profile"); + private static final Instant FIRST = Instant.parse("2024-03-01T00:00:00Z"); + private static final Instant SECOND = Instant.parse("2024-06-01T00:00:00Z"); + private static final TimeInterval VALID = TimeInterval.unboundedFrom( + Instant.parse("2024-01-01T00:00:00Z")); + + private final EventNormalizer normalizer = new EventNormalizer(); + + @Test + public void testCanonicalOrderDoesNotDependOnInputContainers() { + FactKey account = key("account"); + Evidence first = evidence("evidence-a"); + Evidence second = evidence("evidence-b"); + NormalizedMemoryEvent eventA = add("event-a", KEY, FIRST); + NormalizedMemoryEvent eventB = add("event-b", account, FIRST); + NormalizedMemoryEvent eventC = add("event-0-later", KEY, SECOND); + MemoryFactVersion versionA = version( + "opaque-a", FIRST, Arrays.asList(first, first, second)); + MemoryFactVersion versionB = version( + "opaque-b", FIRST, Collections.singletonList(first)); + MemoryFactVersion versionC = version( + "opaque-c", SECOND, Collections.singletonList(second)); + VersionRelation supersedes = new VersionRelation( + VersionRelationType.SUPERSEDES, "opaque-c", "opaque-a"); + VersionRelation duplicate = new VersionRelation( + VersionRelationType.DUPLICATE_OF, "opaque-c", "opaque-a"); + Map> shuffled = new LinkedHashMap<>(); + shuffled.put(KEY, Arrays.asList( + versionC, + version("opaque-a", FIRST, Arrays.asList(second, first, first)))); + shuffled.put(account, Collections.singletonList(versionB)); + Map generating = new LinkedHashMap<>(); + generating.put("opaque-c", "event-0-later"); + generating.put("opaque-b", "event-b"); + generating.put("opaque-a", "event-a"); + CanonicalSnapshot actual = new CanonicalSnapshot( + new TemporalState(shuffled, Arrays.asList(duplicate, supersedes)), + Arrays.asList(eventC, eventB, eventA), + generating); + + Map> ordered = new LinkedHashMap<>(); + ordered.put(account, Collections.singletonList(versionB)); + ordered.put(KEY, Arrays.asList(versionA, versionC)); + CanonicalSnapshot expected = new CanonicalSnapshot( + new TemporalState(ordered, Arrays.asList(supersedes, duplicate)), + Arrays.asList(eventA, eventB, eventC), + generating); + + Assertions.assertEquals(expected, actual); + Assertions.assertEquals(expected.hashCode(), actual.hashCode()); + Assertions.assertEquals( + Arrays.asList(eventA, eventB, eventC), actual.getEvents()); + Assertions.assertEquals( + Arrays.asList(account, KEY), + new ArrayList<>(actual.getState().getVersionsByFactKey().keySet())); + Assertions.assertEquals( + Arrays.asList(versionB, versionA, versionC), + actual.getState().getVersions()); + Assertions.assertEquals( + Arrays.asList("opaque-a", "opaque-b", "opaque-c"), + new ArrayList<>(actual.getGeneratingEventIds().keySet())); + Assertions.assertEquals( + Arrays.asList(supersedes, duplicate), actual.getState().getRelations()); + Assertions.assertEquals( + Arrays.asList(second, first, first), + shuffled.get(KEY).get(1).getEvidence()); + } + + @Test + public void testSnapshotDefensivelyCopiesInputsAndIsDeeplyImmutable() { + NormalizedMemoryEvent event = add("event-a", KEY, FIRST); + MemoryFactVersion version = version( + "opaque-a", FIRST, event.getEvidence()); + List events = new ArrayList<>(); + events.add(event); + Map generating = new LinkedHashMap<>(); + generating.put(version.getId(), event.getEventId()); + CanonicalSnapshot snapshot = new CanonicalSnapshot( + state(KEY, Collections.singletonList(version)), events, generating); + + events.clear(); + generating.clear(); + + Assertions.assertEquals(Collections.singletonList(event), snapshot.getEvents()); + Assertions.assertEquals( + Collections.singletonMap("opaque-a", "event-a"), + snapshot.getGeneratingEventIds()); + Assertions.assertThrows( + UnsupportedOperationException.class, () -> snapshot.getEvents().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.getGeneratingEventIds().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.getState().getVersionsByFactKey().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.getState().getVersionsByFactKey().get(KEY).clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.getState().getVersions().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.getState().getRelations().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.getState().getVersions().get(0).getEvidence().clear()); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> snapshot.getEvents().get(0).getEvidence().clear()); + } + + @Test + public void testPreservesTombstoneAndEventsWithoutMaterializedOutput() { + FactKey relationshipKey = new FactKey("person:alice", "lives_in", "profile"); + MemoryFact relationship = MemoryFact.relationship( + "fact-city", new MemoryEntity("person:alice", "person"), + "lives_in", new MemoryEntity("city:beijing", "city")); + Evidence proof = evidence("evidence-original"); + NormalizedMemoryEvent add = normalizer.normalize( + MemoryEvent.add("event-a", relationship, VALID, FIRST, + Arrays.asList(proof, proof)), + relationshipKey); + NormalizedMemoryEvent retract = normalizer.normalize( + MemoryEvent.retract("event-b", "fact-city", VALID, FIRST, + Collections.singletonList(evidence("evidence-retract"))), + relationshipKey); + MemoryFactVersion tombstone = new MemoryFactVersion( + "opaque-tombstone", relationship, MemoryFactVersionStatus.RETRACTED, + VALID, TimeInterval.unboundedFrom(FIRST), retract.getEvidence()); + CanonicalSnapshot snapshot = new CanonicalSnapshot( + state(relationshipKey, Collections.singletonList(tombstone)), + Arrays.asList(retract, add), + Collections.singletonMap("opaque-tombstone", "event-b")); + + Assertions.assertEquals(Arrays.asList(add, retract), snapshot.getEvents()); + Assertions.assertEquals( + Arrays.asList(proof, proof), snapshot.getEvents().get(0).getEvidence()); + Assertions.assertEquals( + add.getPayloadHash(), snapshot.getEvents().get(0).getPayloadHash()); + Assertions.assertEquals( + Collections.singletonList(tombstone), snapshot.getState().getVersions()); + Assertions.assertEquals( + "city:beijing", + snapshot.getState().getVersions().get(0).getFact().getTarget().get().getId()); + Assertions.assertEquals(1, snapshot.getGeneratingEventIds().size()); + Assertions.assertEquals( + "event-b", snapshot.getGeneratingEventIds().get("opaque-tombstone")); + } + + @Test + public void testEmptyPartitionsAreOmittedAndNullInputsAreRejected() { + TemporalState empty = state(KEY, Collections.emptyList()); + CanonicalSnapshot snapshot = new CanonicalSnapshot( + empty, Collections.emptyList(), Collections.emptyMap()); + + Assertions.assertTrue(snapshot.getState().getVersionsByFactKey().isEmpty()); + Assertions.assertTrue(snapshot.getState().getRelations().isEmpty()); + Assertions.assertTrue(snapshot.getEvents().isEmpty()); + Assertions.assertTrue(snapshot.getGeneratingEventIds().isEmpty()); + Assertions.assertThrows(NullPointerException.class, + () -> new CanonicalSnapshot(null, Collections.emptyList(), Collections.emptyMap())); + Assertions.assertThrows(NullPointerException.class, + () -> new CanonicalSnapshot(empty, null, Collections.emptyMap())); + Assertions.assertThrows(NullPointerException.class, + () -> new CanonicalSnapshot(empty, Collections.emptyList(), null)); + Assertions.assertThrows(NullPointerException.class, + () -> new CanonicalSnapshot( + empty, Collections.singletonList(null), Collections.emptyMap())); + } + + @Test + public void testRepeatedEventIsNoopButReusedEventIdIsRejected() { + NormalizedMemoryEvent original = add("event-a", KEY, FIRST); + NormalizedMemoryEvent reused = add("event-a", KEY, SECOND); + TemporalState empty = state(KEY, Collections.emptyList()); + CanonicalSnapshot once = new CanonicalSnapshot( + empty, Collections.singletonList(original), Collections.emptyMap()); + CanonicalSnapshot repeated = new CanonicalSnapshot( + empty, Arrays.asList(original, original), Collections.emptyMap()); + + Assertions.assertEquals(once, repeated); + Assertions.assertEquals(Collections.singletonList(original), repeated.getEvents()); + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot( + empty, Arrays.asList(original, reused), Collections.emptyMap())); + } + + @Test + public void testRejectsDuplicateVersionIdsAndMismatchedFactKey() { + NormalizedMemoryEvent event = add("event-a", KEY, FIRST); + MemoryFactVersion version = version("opaque-a", FIRST, event.getEvidence()); + List events = Collections.singletonList(event); + Map generating = Collections.singletonMap("opaque-a", "event-a"); + + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot( + state(KEY, Arrays.asList(version, version)), events, generating)); + + Map> duplicateAcrossKeys = new LinkedHashMap<>(); + duplicateAcrossKeys.put(KEY, Collections.singletonList(version)); + duplicateAcrossKeys.put(key("account"), Collections.singletonList(version)); + IllegalArgumentException duplicateId = Assertions.assertThrows( + IllegalArgumentException.class, + () -> new CanonicalSnapshot( + new TemporalState(duplicateAcrossKeys, Collections.emptyList()), + events, generating)); + Assertions.assertEquals("Duplicate version id: opaque-a", duplicateId.getMessage()); + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot( + state(new FactKey("person:bob", "city", "profile"), + Collections.singletonList(version)), + events, generating)); + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot( + state(new FactKey("person:alice", "name", "profile"), + Collections.singletonList(version)), + events, generating)); + } + + @Test + public void testRejectsMissingUnknownOrCrossScopeGenerationReferences() { + NormalizedMemoryEvent event = add("event-a", KEY, FIRST); + MemoryFactVersion version = version("opaque-a", FIRST, event.getEvidence()); + TemporalState state = state(KEY, Collections.singletonList(version)); + List events = Collections.singletonList(event); + + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot(state, events, Collections.emptyMap())); + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot(state, events, + Collections.singletonMap("opaque-a", "event-missing"))); + Map extra = new LinkedHashMap<>(); + extra.put("opaque-a", "event-a"); + extra.put("opaque-missing", "event-a"); + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot(state, events, extra)); + NormalizedMemoryEvent otherScope = add("event-account", key("account"), FIRST); + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot(state, Arrays.asList(event, otherScope), + Collections.singletonMap("opaque-a", "event-account"))); + Assertions.assertThrows(NullPointerException.class, + () -> new CanonicalSnapshot(state, events, + Collections.singletonMap("opaque-a", null))); + } + + @Test + public void testRejectsEitherDanglingRelationEndpoint() { + NormalizedMemoryEvent event = add("event-a", KEY, FIRST); + MemoryFactVersion version = version("opaque-a", FIRST, event.getEvidence()); + Map> versions = Collections.singletonMap( + KEY, Collections.singletonList(version)); + List events = Collections.singletonList(event); + Map generating = Collections.singletonMap("opaque-a", "event-a"); + VersionRelation missingTarget = new VersionRelation( + VersionRelationType.SUPERSEDES, "opaque-a", "opaque-missing"); + VersionRelation missingSource = new VersionRelation( + VersionRelationType.SUPERSEDES, "opaque-missing", "opaque-a"); + + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot( + new TemporalState(versions, Collections.singletonList(missingTarget)), + events, generating)); + Assertions.assertThrows(IllegalArgumentException.class, + () -> new CanonicalSnapshot( + new TemporalState(versions, Collections.singletonList(missingSource)), + events, generating)); + } + + private NormalizedMemoryEvent add(String eventId, FactKey key, Instant recordedAt) { + return normalizer.normalize( + MemoryEvent.add(eventId, fact(), VALID, recordedAt, + Collections.singletonList(evidence("evidence-" + eventId))), + key); + } + + private static MemoryFactVersion version( + String id, Instant recordedAt, List evidence) { + return new MemoryFactVersion( + id, fact(), VALID, TimeInterval.unboundedFrom(recordedAt), evidence); + } + + private static MemoryFact fact() { + return MemoryFact.attribute( + "fact-city", new MemoryEntity("person:alice", "person"), "city", "Beijing"); + } + + private static FactKey key(String scope) { + return new FactKey("person:alice", "city", scope); + } + + private static Evidence evidence(String id) { + return new Evidence(id, new Source("source-registry", "registry"), "Proof " + id); + } + + private static TemporalState state(FactKey key, List versions) { + return new TemporalState(Collections.singletonMap(key, versions), Collections.emptyList()); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/EventLedgerTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/EventLedgerTest.java new file mode 100644 index 000000000..d13f21874 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/EventLedgerTest.java @@ -0,0 +1,148 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.time.Instant; +import java.util.Collections; +import java.util.Optional; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class EventLedgerTest { + + private final EventNormalizer normalizer = new EventNormalizer(); + + @Test + public void testCheckDoesNotMutateBeforeCommit() { + EventLedger ledger = new EventLedger(); + NormalizedMemoryEvent original = normalized( + "event-1", + "Beijing"); + NormalizedMemoryEvent reusedBeforeCommit = normalized( + "event-1", + "Shanghai"); + + Assertions.assertEquals( + EventLedgerDecision.ACCEPTED, + ledger.check(original)); + Assertions.assertEquals( + EventLedgerDecision.ACCEPTED, + ledger.check(reusedBeforeCommit)); + Assertions.assertEquals( + Optional.empty(), + ledger.getPayloadHash("event-1")); + + Assertions.assertEquals( + EventLedgerDecision.ACCEPTED, + ledger.commit(original)); + Assertions.assertEquals( + Optional.of(original.getPayloadHash()), + ledger.getPayloadHash("event-1")); + } + + @Test + public void testDuplicateNoopAndRejectedReuseRemainAtomic() { + EventLedger ledger = new EventLedger(); + NormalizedMemoryEvent original = normalized( + "event-1", + "Beijing"); + NormalizedMemoryEvent duplicate = normalized( + "event-1", + "Beijing"); + NormalizedMemoryEvent reused = normalized( + "event-1", + "Shanghai"); + NormalizedMemoryEvent otherId = normalized( + "event-2", + "Beijing"); + + ledger.commit(original); + + Assertions.assertEquals( + EventLedgerDecision.DUPLICATE_NOOP, + ledger.check(duplicate)); + Assertions.assertEquals( + EventLedgerDecision.DUPLICATE_NOOP, + ledger.commit(duplicate)); + Assertions.assertEquals( + EventLedgerDecision.REJECT_EVENT_ID_REUSE, + ledger.check(reused)); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> ledger.commit(reused)); + Assertions.assertEquals( + Optional.of(original.getPayloadHash()), + ledger.getPayloadHash("event-1")); + Assertions.assertEquals( + EventLedgerDecision.ACCEPTED, + ledger.check(otherId)); + Assertions.assertEquals( + Optional.empty(), + ledger.getPayloadHash("event-2")); + } + + @Test + public void testRejectNullEvent() { + EventLedger ledger = new EventLedger(); + + Assertions.assertThrows( + NullPointerException.class, + () -> ledger.check(null)); + Assertions.assertThrows( + NullPointerException.class, + () -> ledger.commit(null)); + } + + private NormalizedMemoryEvent normalized( + String eventId, + String value) { + MemoryEvent event = MemoryEvent.add( + eventId, + MemoryFact.attribute( + "fact-location", + new MemoryEntity("person:alice", "person"), + "location", + value), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time("2024-03-01T00:00:00Z"), + Collections.singletonList(new Evidence( + "evidence-1", + new Source("source-1", "registry"), + "recorded location"))); + return normalizer.normalize( + event, + new FactKey( + "person:alice", + "location", + "profile")); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/EventNormalizerTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/EventNormalizerTest.java new file mode 100644 index 000000000..4cb044a45 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/EventNormalizerTest.java @@ -0,0 +1,280 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.time.Instant; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; +import java.util.Optional; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.FactValue; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryEventOperation; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class EventNormalizerTest { + + private final EventNormalizer normalizer = new EventNormalizer(); + + @Test + public void testNormalizeUnicodeTimeEvidenceAndEntityReference() { + String decomposedEventId = "event-e\u0301"; + MemoryFact relationship = MemoryFact.relationship( + "fact-location", + new MemoryEntity("person:cafe\u0301", "pe\u0301rson"), + "li\u0301ves_in", + new MemoryEntity("city:be\u0301ijing", "ci\u0301ty")); + MemoryEvent event = MemoryEvent.add( + decomposedEventId, + relationship, + new TimeInterval( + time("2024-01-01T00:00:00.123456789Z"), + time("2025-01-01T00:00:00.987654321Z")), + time("2024-03-01T00:00:00.456789123Z"), + Arrays.asList( + evidence( + "evidence-b", + "source-b", + "registre\u0301-b", + "second"), + evidence( + "evidence-a", + "source-a", + "registre\u0301-a", + "first"))); + FactKey key = new FactKey( + "person:cafe\u0301", + "li\u0301ves_in", + "pro\u0301file"); + + NormalizedMemoryEvent normalized = + normalizer.normalize(event, key); + + Assertions.assertEquals("event-\u00e9", normalized.getEventId()); + Assertions.assertEquals( + MemoryEventOperation.ADD, + normalized.getOperation()); + Assertions.assertEquals("fact-location", normalized.getFactId()); + Assertions.assertEquals( + new FactKey( + "person:caf\u00e9", + "l\u00edves_in", + "pr\u00f3file"), + normalized.getFactKey()); + Assertions.assertEquals( + Optional.of(FactValue.entityReference("city:b\u00e9ijing")), + normalized.getFactValue()); + Assertions.assertEquals( + time("2024-01-01T00:00:00.123456789Z"), + normalized.getValidTime().getStart()); + Assertions.assertEquals( + Optional.of(time("2025-01-01T00:00:00.987654321Z")), + normalized.getValidTime().getEnd()); + Assertions.assertEquals( + time("2024-03-01T00:00:00.456789123Z"), + normalized.getRecordedAt()); + Assertions.assertEquals( + Arrays.asList("evidence-a", "evidence-b"), + Arrays.asList( + normalized.getEvidence().get(0).getId(), + normalized.getEvidence().get(1).getId())); + Assertions.assertEquals( + "registr\u00e9-a", + normalized.getEvidence().get(0).getSource().getName()); + Assertions.assertTrue( + normalized.getPayloadHash().matches("[0-9a-f]{64}")); + Assertions.assertThrows( + UnsupportedOperationException.class, + () -> normalized.getEvidence().clear()); + Assertions.assertEquals(decomposedEventId, event.getId()); + } + + @Test + public void testPayloadHashIsCanonicalAndCoversPayload() { + FactKey key = key("profile"); + Evidence first = evidence( + "evidence-a", + "source-a", + "registry-a", + "first"); + Evidence second = evidence( + "evidence-b", + "source-b", + "registry-b", + "second"); + NormalizedMemoryEvent original = normalizer.normalize( + addEvent( + "event-1", + "Beijing", + "2024-03-01T00:00:00Z", + Arrays.asList(second, first)), + key); + NormalizedMemoryEvent samePayload = normalizer.normalize( + addEvent( + "event-2", + "Beijing", + "2024-03-01T00:00:00Z", + Arrays.asList(first, second)), + key); + NormalizedMemoryEvent differentValue = normalizer.normalize( + addEvent( + "event-1", + "Shanghai", + "2024-03-01T00:00:00Z", + Arrays.asList(first, second)), + key); + NormalizedMemoryEvent differentRecordedAt = normalizer.normalize( + addEvent( + "event-1", + "Beijing", + "2024-04-01T00:00:00Z", + Arrays.asList(first, second)), + key); + NormalizedMemoryEvent differentRecordedAtNanos = normalizer.normalize( + addEvent( + "event-1", + "Beijing", + "2024-03-01T00:00:00.000000001Z", + Arrays.asList(first, second)), + key); + NormalizedMemoryEvent differentScope = normalizer.normalize( + addEvent( + "event-1", + "Beijing", + "2024-03-01T00:00:00Z", + Arrays.asList(first, second)), + key("account")); + + Assertions.assertEquals( + Optional.of(FactValue.literal("Beijing")), + original.getFactValue()); + Assertions.assertEquals( + original.getPayloadHash(), + samePayload.getPayloadHash()); + Assertions.assertNotEquals( + original.getPayloadHash(), + differentValue.getPayloadHash()); + Assertions.assertNotEquals( + original.getPayloadHash(), + differentRecordedAt.getPayloadHash()); + Assertions.assertNotEquals( + original.getPayloadHash(), + differentRecordedAtNanos.getPayloadHash()); + Assertions.assertNotEquals( + original.getPayloadHash(), + differentScope.getPayloadHash()); + } + + @Test + public void testNormalizeRetractAndRejectMismatchedKey() { + FactKey key = key("profile"); + MemoryEvent retract = MemoryEvent.retract( + "event-retract", + "fact-location", + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00.123456789Z")), + time("2024-06-01T00:00:00.987654321Z"), + Collections.singletonList(evidence( + "evidence-a", + "source-a", + "registry-a", + "withdrawn"))); + + NormalizedMemoryEvent normalized = + normalizer.normalize(retract, key); + + Assertions.assertEquals(key, normalized.getFactKey()); + Assertions.assertEquals( + Optional.empty(), + normalized.getFactValue()); + Assertions.assertEquals( + time("2024-06-01T00:00:00.987654321Z"), + normalized.getRecordedAt()); + Assertions.assertThrows( + IllegalArgumentException.class, + () -> normalizer.normalize( + addEvent( + "event-add", + "Beijing", + "2024-03-01T00:00:00Z", + Collections.singletonList(evidence( + "evidence-a", + "source-a", + "registry-a", + "first"))), + new FactKey( + "person:bob", + "location", + "profile"))); + Assertions.assertThrows( + NullPointerException.class, + () -> normalizer.normalize(null, key)); + Assertions.assertThrows( + NullPointerException.class, + () -> normalizer.normalize(retract, null)); + } + + private static MemoryEvent addEvent( + String eventId, + String value, + String recordedAt, + List evidence) { + return MemoryEvent.add( + eventId, + MemoryFact.attribute( + "fact-location", + new MemoryEntity("person:alice", "person"), + "location", + value), + TimeInterval.unboundedFrom( + time("2024-01-01T00:00:00Z")), + time(recordedAt), + evidence); + } + + private static FactKey key(String scope) { + return new FactKey( + "person:alice", + "location", + scope); + } + + private static Evidence evidence( + String evidenceId, + String sourceId, + String sourceName, + String content) { + return new Evidence( + evidenceId, + new Source(sourceId, sourceName), + content); + } + + private static Instant time(String value) { + return Instant.parse(value); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/SnapshotComparatorTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/SnapshotComparatorTest.java new file mode 100644 index 000000000..42cf844e0 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/semantics/SnapshotComparatorTest.java @@ -0,0 +1,259 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.semantics; + +import java.time.Instant; +import java.util.Arrays; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import org.apache.geaflow.ai.temporal.model.Evidence; +import org.apache.geaflow.ai.temporal.model.FactKey; +import org.apache.geaflow.ai.temporal.model.MemoryEntity; +import org.apache.geaflow.ai.temporal.model.MemoryEvent; +import org.apache.geaflow.ai.temporal.model.MemoryFact; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersion; +import org.apache.geaflow.ai.temporal.model.MemoryFactVersionStatus; +import org.apache.geaflow.ai.temporal.model.Source; +import org.apache.geaflow.ai.temporal.model.TimeInterval; +import org.apache.geaflow.ai.temporal.model.VersionRelation; +import org.apache.geaflow.ai.temporal.model.VersionRelationType; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +/** + * Compares events, versions, generators and relations in that order. + * Paths align events and versions by ID; times use ISO-8601 or "infinity". + * Context comes from the expected side, or the present side for a missing record. + * Generator differences name the expected event; relation differences use the + * expected from-version's FactKey and generating event. + */ +public class SnapshotComparatorTest { + + private static final FactKey KEY = new FactKey("person:alice", "location", "profile"); + private static final List EVIDENCE = Collections.singletonList( + new Evidence("evidence-a", new Source("source-a", "registry"), "recorded location")); + + private final SnapshotComparator comparator = new SnapshotComparator(); + + @Test + public void testEquivalentSnapshotsHaveNoDifferenceOrContext() { + CanonicalSnapshot empty = snapshot(Collections.emptyList(), Collections.emptyList(), + Collections.emptyMap(), Collections.emptyList()); + assertEquivalent(comparator.compare(empty, empty)); + + MemoryFactVersion first = version("version-a"); + MemoryFactVersion second = version("version-b"); + NormalizedMemoryEvent firstEvent = event("event-a", "profile"); + NormalizedMemoryEvent secondEvent = event("event-b", "profile"); + CanonicalSnapshot expected = snapshot(Arrays.asList(first, second), + Arrays.asList(firstEvent, secondEvent), generatingIds(), Collections.emptyList()); + Map reversedIds = new LinkedHashMap<>(); + reversedIds.put("version-b", "event-b"); + reversedIds.put("version-a", "event-a"); + CanonicalSnapshot actual = snapshot(Arrays.asList(second, first), + Arrays.asList(secondEvent, firstEvent), reversedIds, Collections.emptyList()); + + assertEquivalent(comparator.compare(expected, actual)); + } + + @Test + public void testStatusPrecedesTransactionEndAndReportsVersionContext() { + CanonicalSnapshot expected = singleVersion(version("version-a")); + CanonicalSnapshot actual = singleVersion(version("version-a", "Paris", + MemoryFactVersionStatus.RETRACTED, Instant.ofEpochMilli(5000))); + + DiffReport difference = comparator.compare(expected, actual); + assertDifference(difference, "versions[version-a].status", + "ACTIVE", "RETRACTED", KEY, "event-a"); + assertEquivalent(comparator.compare(expected, expected)); + assertDifference(difference, "versions[version-a].status", + "ACTIVE", "RETRACTED", KEY, "event-a"); + } + + @Test + public void testTransactionEndUsesExactInstantAndExplicitInfinity() { + CanonicalSnapshot expected = singleVersion(version("version-a")); + CanonicalSnapshot actual = singleVersion(version("version-a", "Paris", + MemoryFactVersionStatus.ACTIVE, Instant.ofEpochMilli(5000))); + + assertDifference(comparator.compare(expected, actual), + "versions[version-a].transactionTime.end", "infinity", + "1970-01-01T00:00:05Z", KEY, "event-a"); + } + + @Test + public void testTransactionEndPreservesNanosecondPrecision() { + Instant expectedEnd = Instant.ofEpochSecond(5, 100); + Instant actualEnd = Instant.ofEpochSecond(5, 200); + CanonicalSnapshot expected = singleVersion(version("version-a", "Paris", + MemoryFactVersionStatus.ACTIVE, expectedEnd)); + CanonicalSnapshot actual = singleVersion(version("version-a", "Paris", + MemoryFactVersionStatus.ACTIVE, actualEnd)); + + assertDifference(comparator.compare(expected, actual), + "versions[version-a].transactionTime.end", + expectedEnd.toString(), actualEnd.toString(), KEY, "event-a"); + } + + @Test + public void testReportsOutputFactValueEvenWhenInputEventsMatch() { + CanonicalSnapshot expected = singleVersion(version("version-a")); + CanonicalSnapshot actual = singleVersion(version("version-a", "Rome", + MemoryFactVersionStatus.ACTIVE, null)); + + assertDifference(comparator.compare(expected, actual), + "versions[version-a].fact.literalValue", "Paris", "Rome", KEY, "event-a"); + } + + @Test + public void testEventOnlyScopeDifferencePrecedesPayloadHash() { + NormalizedMemoryEvent expectedEvent = event("event-a", "profile"); + NormalizedMemoryEvent actualEvent = event("event-a", "private"); + Assertions.assertNotEquals(expectedEvent.getPayloadHash(), actualEvent.getPayloadHash()); + CanonicalSnapshot expected = snapshot(Collections.emptyList(), + Collections.singletonList(expectedEvent), Collections.emptyMap(), + Collections.emptyList()); + CanonicalSnapshot actual = snapshot(Collections.emptyList(), + Collections.singletonList(actualEvent), Collections.emptyMap(), + Collections.emptyList()); + + assertDifference(comparator.compare(expected, actual), "events[event-a].factKey.scope", + "profile", "private", KEY, "event-a"); + } + + @Test + public void testMissingVersionAlignsByIdAndUsesPresentSideContext() { + MemoryFactVersion first = version("version-a"); + MemoryFactVersion second = version("version-b"); + List events = Arrays.asList(event("event-a", "profile"), + event("event-b", "profile")); + CanonicalSnapshot expected = snapshot(Arrays.asList(second, first), events, + generatingIds(), Collections.emptyList()); + CanonicalSnapshot actual = snapshot(Collections.singletonList(second), events, + Collections.singletonMap("version-b", "event-b"), Collections.emptyList()); + + assertDifference(comparator.compare(expected, actual), "versions[version-a]", + "present", null, KEY, "event-a"); + assertDifference(comparator.compare(actual, expected), "versions[version-a]", + null, "present", KEY, "event-a"); + } + + @Test + public void testGeneratingEventDifferenceUsesExpectedEventContext() { + List versions = Collections.singletonList(version("version-a")); + List events = Arrays.asList(event("event-a", "profile"), + event("event-b", "profile")); + CanonicalSnapshot expected = snapshot(versions, events, + Collections.singletonMap("version-a", "event-a"), Collections.emptyList()); + CanonicalSnapshot actual = snapshot(versions, events, + Collections.singletonMap("version-a", "event-b"), Collections.emptyList()); + + assertDifference(comparator.compare(expected, actual), "generatingEventIds[version-a]", + "event-a", "event-b", KEY, "event-a"); + } + + @Test + public void testRelationTypeDifferenceReportsFromVersionContext() { + List versions = Arrays.asList(version("version-a"), + version("version-b")); + List events = Arrays.asList(event("event-a", "profile"), + event("event-b", "profile")); + CanonicalSnapshot expected = snapshot(versions, events, generatingIds(), + Collections.singletonList(new VersionRelation(VersionRelationType.SUPERSEDES, + "version-b", "version-a"))); + CanonicalSnapshot actual = snapshot(versions, events, generatingIds(), + Collections.singletonList(new VersionRelation(VersionRelationType.DUPLICATE_OF, + "version-b", "version-a"))); + + assertDifference(comparator.compare(expected, actual), "relations[0].type", + "SUPERSEDES", "DUPLICATE_OF", KEY, "event-b"); + } + + private static void assertEquivalent(DiffReport report) { + Assertions.assertTrue(report.isEquivalent()); + Assertions.assertEquals(Optional.empty(), report.getFieldPath()); + Assertions.assertEquals(Optional.empty(), report.getExpectedValue()); + Assertions.assertEquals(Optional.empty(), report.getActualValue()); + Assertions.assertEquals(Optional.empty(), report.getFactKey()); + Assertions.assertEquals(Optional.empty(), report.getEventId()); + } + + private static void assertDifference(DiffReport report, String path, String expected, + String actual, FactKey key, String eventId) { + Assertions.assertFalse(report.isEquivalent()); + Assertions.assertEquals(Optional.of(path), report.getFieldPath()); + Assertions.assertEquals(Optional.ofNullable(expected), report.getExpectedValue()); + Assertions.assertEquals(Optional.ofNullable(actual), report.getActualValue()); + Assertions.assertEquals(Optional.of(key), report.getFactKey()); + Assertions.assertEquals(Optional.of(eventId), report.getEventId()); + } + + private static CanonicalSnapshot singleVersion(MemoryFactVersion version) { + return snapshot(Collections.singletonList(version), + Collections.singletonList(event("event-a", "profile")), + Collections.singletonMap(version.getId(), "event-a"), Collections.emptyList()); + } + + private static CanonicalSnapshot snapshot(List versions, + List events, + Map generatingEventIds, + List relations) { + Map> byKey = versions.isEmpty() + ? Collections.emptyMap() : Collections.singletonMap(KEY, versions); + return new CanonicalSnapshot(new TemporalState(byKey, relations), events, + generatingEventIds); + } + + private static Map generatingIds() { + Map ids = new LinkedHashMap<>(); + ids.put("version-a", "event-a"); + ids.put("version-b", "event-b"); + return ids; + } + + private static NormalizedMemoryEvent event(String eventId, String scope) { + MemoryEvent event = MemoryEvent.add(eventId, fact("Paris"), validTime(), + Instant.ofEpochMilli(2000), EVIDENCE); + return new EventNormalizer().normalize(event, + new FactKey(KEY.getSubjectId(), KEY.getPredicate(), scope)); + } + + private static MemoryFactVersion version(String versionId) { + return version(versionId, "Paris", MemoryFactVersionStatus.ACTIVE, null); + } + + private static MemoryFactVersion version(String versionId, String value, + MemoryFactVersionStatus status, Instant end) { + return new MemoryFactVersion(versionId, fact(value), status, validTime(), + new TimeInterval(Instant.ofEpochMilli(2000), end), EVIDENCE); + } + + private static MemoryFact fact(String value) { + return MemoryFact.attribute("fact-location", new MemoryEntity("person:alice", "person"), + "location", value); + } + + private static TimeInterval validTime() { + return new TimeInterval(Instant.ofEpochMilli(1000), Instant.ofEpochMilli(9000)); + } +} diff --git a/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/udga/TemporalUdgaFeasibilityTest.java b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/udga/TemporalUdgaFeasibilityTest.java new file mode 100644 index 000000000..277d19b43 --- /dev/null +++ b/geaflow-ai/src/test/java/org/apache/geaflow/ai/temporal/udga/TemporalUdgaFeasibilityTest.java @@ -0,0 +1,230 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.geaflow.ai.temporal.udga; + +import java.io.IOException; +import java.io.Serializable; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.HashMap; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.stream.Stream; +import org.apache.geaflow.cluster.system.ClusterMetaStore; +import org.apache.geaflow.common.config.Configuration; +import org.apache.geaflow.common.config.keys.DSLConfigKeys; +import org.apache.geaflow.common.config.keys.ExecutionConfigKeys; +import org.apache.geaflow.dsl.connector.file.FileConstants; +import org.apache.geaflow.dsl.runtime.QueryClient; +import org.apache.geaflow.dsl.runtime.QueryContext; +import org.apache.geaflow.dsl.runtime.engine.GQLPipeLine; +import org.apache.geaflow.dsl.runtime.engine.GQLPipeLine.GQLPipelineHook; +import org.apache.geaflow.env.Environment; +import org.apache.geaflow.env.EnvironmentFactory; +import org.apache.geaflow.file.FileConfigKeys; +import org.apache.geaflow.runtime.core.scheduler.resource.ScheduledWorkerManagerFactory; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +public class TemporalUdgaFeasibilityTest { + + private static final String QUERY_RESOURCE = + "/temporal/temporal_udga_feasibility.sql"; + + @TempDir + Path tempDirectory; + + @Test + public void testRowUdgaCarriesStateAcrossTwoDynamicEdgeBatches() + throws Exception { + Path vertices = writeInput( + "vertices.csv", + "1,Alice", + "2,Bob", + "3,Carol"); + Path edges = writeInput( + "edges.csv", + "1,2", + "2,3"); + Path output = tempDirectory.resolve("result"); + Environment environment = null; + + try { + environment = EnvironmentFactory.onLocalEnvironment(); + environment.getEnvironmentContext().withConfig( + localConfiguration()); + GQLPipeLine pipeline = new GQLPipeLine(environment, 0); + pipeline.setPipelineHook(new PathReplacingHook( + vertices, + edges, + output)); + + pipeline.execute(); + + List results = readResults(output); + Set observedBatchCounts = new HashSet<>(); + for (ProbeResult result : results) { + observedBatchCounts.add(result.batchCount); + } + Assertions.assertEquals( + new HashSet<>(Arrays.asList(1, 2)), + observedBatchCounts); + Assertions.assertTrue(results.stream().anyMatch(result -> + result.vertexId == 2L + && result.batchCount == 2 + && result.hadPreviousValue)); + Assertions.assertTrue(results.stream().anyMatch(result -> + result.dynamicEdgeCount > 0)); + Assertions.assertTrue(results.stream().anyMatch(result -> + result.receivedMessageCount > 0)); + } finally { + if (environment != null) { + environment.shutdown(); + } + ClusterMetaStore.close(); + ScheduledWorkerManagerFactory.clear(); + } + } + + private Path writeInput(String fileName, String... lines) + throws IOException { + Path path = tempDirectory.resolve(fileName); + Files.write(path, Arrays.asList(lines), StandardCharsets.UTF_8); + return path; + } + + private Map localConfiguration() { + Map config = new HashMap<>(); + config.put( + DSLConfigKeys.GEAFLOW_DSL_QUERY_PATH.getKey(), + FileConstants.PREFIX_JAVA_RESOURCE + QUERY_RESOURCE); + config.put( + ExecutionConfigKeys.JOB_APP_NAME.getKey(), + "TemporalUdgaFeasibilityTest"); + config.put( + ExecutionConfigKeys.JOB_WORK_PATH.getKey(), + tempDirectory.resolve("work").toString()); + config.put( + FileConfigKeys.ROOT.getKey(), + tempDirectory.resolve("state").toString()); + return config; + } + + private static List readResults(Path output) + throws IOException { + List results = new ArrayList<>(); + try (Stream paths = Files.walk(output)) { + for (Path path : (Iterable) paths + .filter(Files::isRegularFile)::iterator) { + for (String line : Files.readAllLines( + path, + StandardCharsets.UTF_8)) { + if (!line.trim().isEmpty()) { + results.add(ProbeResult.parse(line)); + } + } + } + } + Assertions.assertFalse(results.isEmpty()); + return results; + } + + private static String sqlPath(Path path) { + return path.toAbsolutePath().toString().replace('\\', '/'); + } + + private static final class PathReplacingHook + implements GQLPipelineHook, Serializable { + + private final String vertices; + private final String edges; + private final String output; + + private PathReplacingHook( + Path vertices, + Path edges, + Path output) { + this.vertices = sqlPath(vertices); + this.edges = sqlPath(edges); + this.output = sqlPath(output); + } + + @Override + public String rewriteScript( + String script, + Configuration configuration) { + return script + .replace("${vertices}", vertices) + .replace("${edges}", edges) + .replace("${output}", output); + } + + @Override + public void beforeExecute( + QueryClient queryClient, + QueryContext queryContext) { + } + + @Override + public void afterExecute( + QueryClient queryClient, + QueryContext queryContext) { + } + } + + private static final class ProbeResult { + + private final long vertexId; + private final int batchCount; + private final boolean hadPreviousValue; + private final int dynamicEdgeCount; + private final int receivedMessageCount; + + private ProbeResult( + long vertexId, + int batchCount, + boolean hadPreviousValue, + int dynamicEdgeCount, + int receivedMessageCount) { + this.vertexId = vertexId; + this.batchCount = batchCount; + this.hadPreviousValue = hadPreviousValue; + this.dynamicEdgeCount = dynamicEdgeCount; + this.receivedMessageCount = receivedMessageCount; + } + + private static ProbeResult parse(String line) { + String[] fields = line.split(","); + Assertions.assertEquals(5, fields.length, line); + return new ProbeResult( + Long.parseLong(fields[0]), + Integer.parseInt(fields[1]), + Boolean.parseBoolean(fields[2]), + Integer.parseInt(fields[3]), + Integer.parseInt(fields[4])); + } + } +} diff --git a/geaflow-ai/src/test/resources/temporal/temporal_udga_feasibility.sql b/geaflow-ai/src/test/resources/temporal/temporal_udga_feasibility.sql new file mode 100644 index 000000000..9e9b9f833 --- /dev/null +++ b/geaflow-ai/src/test/resources/temporal/temporal_udga_feasibility.sql @@ -0,0 +1,87 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +CREATE FUNCTION temporal_udga_probe AS +'org.apache.geaflow.ai.temporal.udga.TemporalUdgaProbe'; + +CREATE GRAPH temporal_probe_graph ( + Vertex entity ( + id bigint ID, + name varchar + ), + Edge relates ( + src_id bigint SOURCE ID, + target_id bigint DESTINATION ID + ) +) WITH ( + storeType = 'memory', + shardCount = 1 +); + +CREATE TABLE temporal_probe_vertices ( + id bigint, + name varchar +) WITH ( + type = 'file', + geaflow.dsl.file.path = '${vertices}', + geaflow.dsl.window.size = -1 +); + +CREATE TABLE temporal_probe_edges ( + src_id bigint, + target_id bigint +) WITH ( + type = 'file', + geaflow.dsl.file.path = '${edges}', + geaflow.dsl.window.size = 1 +); + +INSERT INTO temporal_probe_graph.entity +SELECT id, name FROM temporal_probe_vertices; + +INSERT INTO temporal_probe_graph.relates +SELECT src_id, target_id FROM temporal_probe_edges; + +CREATE TABLE temporal_probe_results ( + vertex_id bigint, + batch_count int, + had_previous_value boolean, + dynamic_edge_count int, + received_message_count int +) WITH ( + type = 'file', + geaflow.dsl.file.path = '${output}' +); + +USE GRAPH temporal_probe_graph; + +INSERT INTO temporal_probe_results +CALL temporal_udga_probe() YIELD ( + vertex_id, + batch_count, + had_previous_value, + dynamic_edge_count, + received_message_count +) +RETURN + vertex_id, + batch_count, + had_previous_value, + dynamic_edge_count, + received_message_count;