Skip to content
Open
Show file tree
Hide file tree
Changes from 6 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -27,15 +27,16 @@ import com.embabel.agent.core.Export
import com.embabel.agent.core.support.NIRVANA
import com.embabel.agent.core.support.Rerun
import com.embabel.agent.core.support.safelyGetToolsFrom
import com.embabel.agent.spi.validation.AchievableGoalValidator
import com.embabel.agent.spi.validation.AgentStructureAgentValidator
import com.embabel.agent.spi.validation.DefaultAgentValidationManager
import com.embabel.agent.spi.validation.GoapPathToCompletionValidator
import com.embabel.agent.spi.validation.PathToCompletionAgentValidator
import com.embabel.agent.spi.validation.isActionMethod
import com.embabel.agent.spi.validation.isConditionMethod
import com.embabel.agent.spi.validation.isMethodFromSupertype
import com.embabel.common.core.types.Semver
import com.embabel.common.util.NameUtils
import com.embabel.common.util.loggerFor
import com.fasterxml.jackson.annotation.JsonTypeInfo
import tools.jackson.databind.annotation.JsonDeserialize
import org.slf4j.LoggerFactory
import org.springframework.beans.factory.annotation.Value
import org.springframework.cglib.proxy.Enhancer
Expand Down Expand Up @@ -110,23 +111,20 @@ internal data class AgenticInfo(
class AgentMetadataReader(
private val actionMethodManager: ActionMethodManager = DefaultActionMethodManager(),
private val nameGenerator: MethodDefinedOperationNameGenerator = MethodDefinedOperationNameGenerator(),
agentStructureValidator: AgentStructureAgentValidator = AgentStructureAgentValidator.PERMIT_ALL,
pathToCompletionValidator: PathToCompletionAgentValidator = GoapPathToCompletionValidator(),
private val agentStructureValidator: AgentStructureAgentValidator = AgentStructureAgentValidator.PERMIT_ALL,
private val pathToCompletionValidator: PathToCompletionAgentValidator = GoapPathToCompletionValidator(),
private val requireInterfaceDeserializationAnnotations: Boolean = false,
@Value("\${embabel.agent.platform.planner.restricted-goals:false}")
private val restrictedGoals: Boolean = false,
@Value("\${embabel.agent.api.validation.manager.skip-agent-deployment-on-error:false}")
private val skipAgentDeploymentOnError: Boolean = false,
) {

private val supervisorAgentFactory = SupervisorAgentFactory()

private val logger = LoggerFactory.getLogger(AgentMetadataReader::class.java)

private val agentValidationManager: AgentValidationManager = DefaultAgentValidationManager(
listOf(
agentStructureValidator,
pathToCompletionValidator
)
)
private lateinit var agentValidationManager: AgentValidationManager

fun createAgentScopes(vararg instances: Any): List<AgentScope> =
instances.mapNotNull { createAgentMetadata(it) }
Expand All @@ -152,13 +150,15 @@ class AgentMetadataReader(
val targetType = agenticInfo.getTargetType()

if (!agenticInfo.agentic()) {
logger.debug(
logger.warn(
"No @{} or @{} annotation found on {}",
EmbabelComponent::class.simpleName,
Agent::class.simpleName,
targetType.name,
)
return null
if (skipAgentDeploymentOnError) {

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@igordayen Can you look into EmbabelMockitoIntegrationTestBlockingTest. It's failing because the it's not returning by default.
Looks likeAutoRegistration.kt can pick up any Bean and it's really depend on this check to not behind any condition.
Therefore, I am not putting return null behind property check.

return null
}
}

if (agenticInfo.validationErrors().isNotEmpty()) {
Expand All @@ -168,10 +168,21 @@ class AgentMetadataReader(
Agent::class.simpleName,
targetType.name,
)
return null
if (skipAgentDeploymentOnError) {
return null
}
}
rejectOperationContextConstructorInjection(targetType)

val plannerType = agenticInfo.agentAnnotation?.planner ?: PlannerType.GOAP
agentValidationManager = DefaultAgentValidationManager(
listOf(
agentStructureValidator,
pathToCompletionValidator,
AchievableGoalValidator(agenticInfo.agentName(), targetType, instance, requireInterfaceDeserializationAnnotations)
)
)

val getterGoals = findGoalGetters(targetType).map { getGoal(it, instance) }
val actionMethods = findActionMethods(targetType)
val conditionMethods = findConditionMethods(targetType)
Expand Down Expand Up @@ -204,8 +215,6 @@ class AgentMetadataReader(
)
}

val plannerType = agenticInfo.agentAnnotation?.planner ?: PlannerType.GOAP

val goals = buildSet {
addAll(getterGoals)
addAll(allGoals)
Expand All @@ -222,27 +231,23 @@ class AgentMetadataReader(
Condition::class.simpleName,
targetType.name,
)
return null
if (skipAgentDeploymentOnError) {
return null
}
}

val agent = if (agenticInfo.agentAnnotation != null) {
val goalActions = actionMethods.filter { it.isAnnotationPresent(AchievesGoal::class.java) }
if (plannerType == PlannerType.SUPERVISOR) {
// Find the goal action (the action with @AchievesGoal)
if (goalActions.isEmpty()) {
logger.warn(
"SUPERVISOR planner requires at least one @AchievesGoal action on {}",
targetType.name,
)
return null
}
if (goalActions.size > 1) {
logger.warn(
"SUPERVISOR planner currently supports only one @AchievesGoal action, found {} on {}",
goalActions.size,
targetType.name,
)
return null
if (skipAgentDeploymentOnError) {
return null
}
}
val goalAction = allActions.find { action ->
goalActions.any { method ->
Expand Down Expand Up @@ -304,8 +309,9 @@ class AgentMetadataReader(
val validationResult = agentValidationManager.validate(agent)
if (!validationResult.isValid) {
logger.warn("Agent validation failed:\n${validationResult.errors.joinToString("\n")}")
// TODO: Uncomment to strengthen validation and refactor the test if needed. Because some tests might fail.
// return null
if (skipAgentDeploymentOnError) {
return null
}
}
}

Expand Down Expand Up @@ -340,7 +346,7 @@ class AgentMetadataReader(
stateClass,
)
allActions.add(action)
createGoalFromStateActionMethod(actionMethod, action, stateClass, agentInstance)?.let {
createGoalFromStateActionMethod(actionMethod, action, stateClass)?.let {
allGoals.add(it)
}
// Recursively unroll if this action also returns a @State type
Expand Down Expand Up @@ -431,7 +437,6 @@ class AgentMetadataReader(
method: Method,
action: CoreAction,
stateClass: Class<*>,
agentInstance: Any,
): AgentCoreGoal? {
val actionAnnotation = method.getAnnotation(Action::class.java)
val goalAnnotation = method.getAnnotation(AchievesGoal::class.java) ?: return null
Expand Down Expand Up @@ -511,68 +516,13 @@ class AgentMetadataReader(
type,
{ method -> actionMethods.add(method) },
// Get annotated methods from this type and interfaces
{ method -> isActionMethod(method, type) })
{ method -> isActionMethod(logger,method, type, requireInterfaceDeserializationAnnotations) })
if (actionMethods.isEmpty()) {
logger.debug("No methods annotated with @{} found in {}", Action::class.simpleName, type)
}
return actionMethods
}

private fun isActionMethod(
method: Method,
type: Class<*>,
): Boolean {
return method.isAnnotationPresent(Action::class.java) &&
(type.declaredMethods.contains(method) || isMethodFromSupertype(method, type)) &&
(!method.returnType.isInterface || !requireInterfaceDeserializationAnnotations || hasRequiredJsonDeserializeAnnotationOnInterfaceReturnType(
method
))
}

private fun isConditionMethod(
method: Method,
type: Class<*>,
): Boolean {
return method.isAnnotationPresent(Condition::class.java) &&
(type.declaredMethods.contains(method) || isMethodFromSupertype(method, type))
}

private fun isMethodFromSupertype(
method: Method,
type: Class<*>,
): Boolean {
// Check interfaces
if (type.interfaces.any { interfaceType ->
interfaceType.declaredMethods.any { interfaceMethod ->
methodSignaturesMatch(method, interfaceMethod)
}
}) {
return true
}

// Check superclasses
var superclass = type.superclass
while (superclass != null && superclass != Any::class.java) {
if (superclass.declaredMethods.any { superMethod ->
methodSignaturesMatch(method, superMethod)
}) {
return true
}
superclass = superclass.superclass
}

return false
}

private fun methodSignaturesMatch(
method1: Method,
method2: Method,
): Boolean {
return method1.name == method2.name &&
method1.parameterTypes.contentEquals(method2.parameterTypes) &&
method1.returnType == method2.returnType
}

private fun findGoalGetters(type: Class<*>): List<Method> {
val goalGetters = mutableListOf<Method>()
type.declaredMethods.forEach { method ->
Expand Down Expand Up @@ -781,22 +731,3 @@ private fun rejectOperationContextConstructorInjection(agentClass: Class<*>) {
)
}
}

/**
* Checks if a method returning an interface returns a type with a @JsonDeserialize annotation.
* @param method The Java method to check.
* @return true if the return type has a @JsonDeserialize annotation, false otherwise
*/
private fun hasRequiredJsonDeserializeAnnotationOnInterfaceReturnType(method: Method): Boolean {
val hasRequiredAnnotation = method.returnType.isAnnotationPresent(JsonDeserialize::class.java) ||
method.returnType.isAnnotationPresent(JsonTypeInfo::class.java)
if (!hasRequiredAnnotation) {
loggerFor<AgentMetadataReader>().warn(
"❓Interface {} used as return type of {}.{} must have @JsonDeserialize or @JsonTypeInfo annotation",
method.returnType.name,
method.declaringClass.name,
method.name,
)
}
return hasRequiredAnnotation
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
package com.embabel.agent.spi.validation

import com.embabel.agent.api.annotation.AchievesGoal
import com.embabel.agent.core.AgentScope
import com.embabel.common.core.validation.ValidationError
import com.embabel.common.core.validation.ValidationErrorCodes
import com.embabel.common.core.validation.ValidationLocation
import com.embabel.common.core.validation.ValidationResult
import com.embabel.common.core.validation.ValidationSeverity
import org.slf4j.LoggerFactory
import org.springframework.util.ReflectionUtils
import java.lang.reflect.Method

/**
* Validator that checks methods annotated with AchievesGoal.
*
* Specific check includes:
* - Verifying that @Action annotation is present on it.
*/
open class AchievableGoalValidator ( private val agentName: String,
private val agentClass: Class<*>,
private val agentInstance: Any,
private val requireInterfaceDeserializationAnnotations: Boolean): AgentValidator
{
private val logger = LoggerFactory.getLogger(AchievableGoalValidator::class.java)

private fun isMethodAnnotatedWithAchievesGoal(
method: Method,
): Boolean {
return method.isAnnotationPresent(AchievesGoal::class.java)
}

override fun validate(agentScope: AgentScope): ValidationResult {
val errors = mutableListOf<ValidationError>()
ReflectionUtils.doWithMethods(
agentClass,
{ method ->
if(!isActionMethod(logger,method, agentClass, requireInterfaceDeserializationAnnotations)) {
errors.add(
ValidationError(
code = ValidationErrorCodes.MISSING_ACTION_ANNOTATION,
message = "@Action annotation is missing on the method '${agentInstance.javaClass.name}.${method.name}' annotated with @AchievesGoal.",
severity = ValidationSeverity.ERROR,
location = ValidationLocation(
type = "Agent",
name = agentInstance.javaClass.name,
agentName = agentName,
component = method.name
)
)
)
}
},
{ method -> isMethodAnnotatedWithAchievesGoal(method) })
return ValidationResult(errors.isEmpty(), errors)
}
}
Loading
Loading