diff --git a/src/main/java/org/springframework/persistence/graph/Relationship.java b/src/main/java/org/springframework/persistence/graph/Relationship.java index aa4489ea2..cf6704ec1 100644 --- a/src/main/java/org/springframework/persistence/graph/Relationship.java +++ b/src/main/java/org/springframework/persistence/graph/Relationship.java @@ -5,6 +5,8 @@ import java.lang.annotation.Retention; import java.lang.annotation.RetentionPolicy; import java.lang.annotation.Target; +import org.springframework.persistence.graph.neo4j.NodeBacked; + @Retention(RetentionPolicy.RUNTIME) @Target(ElementType.FIELD) public @interface Relationship { @@ -13,6 +15,6 @@ public @interface Relationship { Direction direction(); - Class elementClass() default Void.class; + Class elementClass() default NodeBacked.class; } diff --git a/src/main/java/org/springframework/persistence/graph/neo4j/Neo4jNodeBacking.aj b/src/main/java/org/springframework/persistence/graph/neo4j/Neo4jNodeBacking.aj index 53ff8cdd7..b0573853b 100644 --- a/src/main/java/org/springframework/persistence/graph/neo4j/Neo4jNodeBacking.aj +++ b/src/main/java/org/springframework/persistence/graph/neo4j/Neo4jNodeBacking.aj @@ -1,10 +1,13 @@ package org.springframework.persistence.graph.neo4j; import java.lang.reflect.Field; +import java.util.AbstractSet; import java.util.Collection; import java.util.HashSet; +import java.util.Iterator; import java.util.Set; +import org.aspectj.lang.ProceedingJoinPoint; import org.aspectj.lang.reflect.FieldSignature; import org.neo4j.graphdb.Direction; import org.neo4j.graphdb.DynamicRelationshipType; @@ -141,7 +144,7 @@ public aspect Neo4jNodeBacking extends AbstractTypeAnnotatingMixinFields Neo4J simple node property [" + propName + "] with value=[" + newVal + "]"); - return null; + return proceed(entity, newVal); } RelationshipInfo relInfo = relationshipInfoFactory.forField(f); @@ -150,8 +153,8 @@ public aspect Neo4jNodeBacking extends AbstractTypeAnnotatingMixinFields Neo4J relationship with value=[" + newVal + "]"); - relInfo.apply(entity, newVal); - return null; + Object result=relInfo.apply(entity, newVal); + return proceed(entity,result); } catch(NotInTransactionException e) { throw new InvalidDataAccessResourceUsageException("Not in a Neo4j transaction.", e); } @@ -168,7 +171,7 @@ public aspect Neo4jNodeBacking extends AbstractTypeAnnotatingMixinFields relatedType=(Class)field.getType(); if (relAnnotation != null) { return new SingleRelationshipInfo(DynamicRelationshipType.withName(relAnnotation.type()), - relAnnotation.direction().toNeo4jDir(), field.getType(), graphEntityInstantiator); + relAnnotation.direction().toNeo4jDir(),relatedType, graphEntityInstantiator); } return new SingleRelationshipInfo(DynamicRelationshipType.withName(getNeo4jPropertyName(field)), - Direction.OUTGOING, field.getType(), graphEntityInstantiator); + Direction.OUTGOING, relatedType, graphEntityInstantiator); } if (isOneToNRelationshipField(field)) { - final Relationship relAnnotation = field.getAnnotation(Relationship.class); return new OneToNRelationshipInfo(DynamicRelationshipType.withName(relAnnotation.type()), relAnnotation.direction().toNeo4jDir(), relAnnotation.elementClass(), graphEntityInstantiator); } @@ -200,14 +203,13 @@ public aspect Neo4jNodeBacking extends AbstractTypeAnnotatingMixinFields clazz; + private final Class relatedType; private final EntityInstantiator graphEntityInstantiator; - public SingleRelationshipInfo(RelationshipType type, Direction direction, Class clazz, EntityInstantiator graphEntityInstantiator) { + public SingleRelationshipInfo(RelationshipType type, Direction direction, Class clazz, EntityInstantiator graphEntityInstantiator) { this.type = type; this.direction = direction; - this.clazz = clazz; + this.relatedType = clazz; this.graphEntityInstantiator = graphEntityInstantiator; } - public void apply(NodeBacked entity, Object newVal) { + public Object apply(NodeBacked entity, Object newVal) { if (newVal != null && !(newVal instanceof NodeBacked)) { throw new IllegalArgumentException("New value must be NodeBacked."); } @@ -234,7 +236,7 @@ public aspect Neo4jNodeBacking extends AbstractTypeAnnotatingMixinFields) clazz); + return graphEntityInstantiator.createEntityFromState(targetNode, relatedType); } } public static class OneToNRelationshipInfo implements RelationshipInfo { + private static final class ManagedSet extends AbstractSet { + private final NodeBacked entity; + final Set delegate; + private final RelationshipInfo relationshipInfo; + + private ManagedSet(NodeBacked entity, Object newVal, RelationshipInfo relationshipInfo) { + this.entity = entity; + this.relationshipInfo = relationshipInfo; + delegate = (Set) newVal; + } + + @Override + public Iterator iterator() { + return delegate.iterator(); + } + + @Override + public int size() { + return delegate.size(); + } + + @Override + public boolean add(NodeBacked e) { + boolean res=delegate.add(e); + if (res) { + relationshipInfo.apply(entity, delegate); + } + return res; + } + } + private final RelationshipType type; private final Direction direction; - private final Class elementClass; + private final Class relatedType; private final EntityInstantiator graphEntityInstantiator; - public OneToNRelationshipInfo(RelationshipType type, Direction direction, Class elementClass, EntityInstantiator graphEntityInstantiator) { + public OneToNRelationshipInfo(RelationshipType type, Direction direction, Class elementClass, EntityInstantiator graphEntityInstantiator) { this.type = type; this.direction = direction; - this.elementClass = elementClass; + this.relatedType = elementClass; this.graphEntityInstantiator = graphEntityInstantiator; } - public void apply(NodeBacked entity, Object newVal) { + public Object apply(final NodeBacked entity, final Object newVal) { Node entityNode = entity.getUnderlyingNode(); + Set newNodes=new HashSet(); if (newVal != null) { if (!(newVal instanceof Set)) { throw new IllegalArgumentException("New value must be a Set, was: " + newVal.getClass()); @@ -294,28 +325,29 @@ public aspect Neo4jNodeBacking extends AbstractTypeAnnotatingMixinFields) newVal) { - NodeBacked nb = (NodeBacked) obj; - Node targetNode = nb.getUnderlyingNode(); + for (Node newNode : newNodes) { switch(direction) { - case OUTGOING : entityNode.createRelationshipTo(targetNode, type); break; - case INCOMING : targetNode.createRelationshipTo(entityNode, type); break; + case OUTGOING : entityNode.createRelationshipTo(newNode, type); break; + case INCOMING : newNode.createRelationshipTo(entityNode, type); break; default : throw new IllegalArgumentException("invalid direction " + direction); } } - + return new ManagedSet(entity, newVal,this); // TODO managedSet that for each mutating method calls this apply (todo use AspectJ to handle that?) } @Override @@ -324,12 +356,12 @@ public aspect Neo4jNodeBacking extends AbstractTypeAnnotatingMixinFields rels = entityNode.getRelationships(type, direction); - Set result = new HashSet(); - for (org.neo4j.graphdb.Relationship rel : rels) { - result.add(graphEntityInstantiator.createEntityFromState(rel.getOtherNode(entityNode), (Class) elementClass)); + Set result = new HashSet(); + for (org.neo4j.graphdb.Relationship rel : entityNode.getRelationships(type, direction)) { + NodeBacked newEntity=graphEntityInstantiator.createEntityFromState(rel.getOtherNode(entityNode), relatedType);; + result.add(newEntity); } - return result; + return new ManagedSet(entity, result,this); } } diff --git a/src/test/java/org/springframework/persistence/test/Person.java b/src/test/java/org/springframework/persistence/test/Person.java index 06eb3bd05..064b0fbc4 100644 --- a/src/test/java/org/springframework/persistence/test/Person.java +++ b/src/test/java/org/springframework/persistence/test/Person.java @@ -93,5 +93,8 @@ public class Person { public void setFriend(Person friend) { this.friend = friend; } - + @Override + public String toString() { + return name; + } } diff --git a/src/test/java/org/springframework/persistence/test/graph/Neo4jGraphPersistenceTest.java b/src/test/java/org/springframework/persistence/test/graph/Neo4jGraphPersistenceTest.java index 2697e9000..4d3178ab3 100644 --- a/src/test/java/org/springframework/persistence/test/graph/Neo4jGraphPersistenceTest.java +++ b/src/test/java/org/springframework/persistence/test/graph/Neo4jGraphPersistenceTest.java @@ -205,6 +205,20 @@ public class Neo4jGraphPersistenceTest { Assert.assertTrue(Set.class.isAssignableFrom(personsFromGet.getClass())); } + @Test + @Transactional + public void testAddToOneToManyRelationship() { + Person michael = new Person("Michael", 35); + Person david = new Person("David", 25); + Group group = new Group(); + group.setPersons(new HashSet()); + group.addPerson(michael); + group.addPerson(david); + Collection personsFromGet = group.getPersons(); + Assert.assertEquals(new HashSet(Arrays.asList(david,michael)), personsFromGet); + Assert.assertTrue(Set.class.isAssignableFrom(personsFromGet.getClass())); + } + @Test @Transactional public void testInstantiatedFinder() {