Skip to content
Draft
Show file tree
Hide file tree
Changes from all 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
11 changes: 8 additions & 3 deletions .github/workflows/ci-dataflow.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -28,17 +28,22 @@ concurrency:

jobs:
check:
name: Run tests on JDK ${{ matrix.jdk }}
runs-on: ubuntu-latest
container: gitlab/gitlab-runner-helper:ubuntu-x86_64-latest
strategy:
fail-fast: false
matrix:
jdk: [ 11, 17 ]
permissions:
contents: read
steps:
- uses: actions/checkout@v4

- name: Set up JDK 11
- name: Set up JDK ${{ matrix.jdk }}
uses: actions/setup-java@v4
with:
java-version: '11'
java-version: ${{ matrix.jdk }}
distribution: 'temurin'

- name: Install Go-ir dependencies
Expand All @@ -63,6 +68,6 @@ jobs:
if: (!cancelled())
uses: actions/upload-artifact@v4
with:
name: gradle-reports-ci-dataflow
name: gradle-reports-ci-dataflow-jdk${{ matrix.jdk }}
path: '**/build/reports/'
retention-days: 1
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,13 @@ plugins {
java
}

val recordSamplePath = "sample/alias/RecordAliasSample.java"
val supportsJava17 = JavaVersion.current().isCompatibleWith(JavaVersion.VERSION_17)

sourceSets.main {
java.exclude(recordSamplePath)
}

tasks {
withType<JavaCompile> {
sourceCompatibility = JavaVersion.VERSION_1_8.toString()
Expand All @@ -10,7 +17,21 @@ tasks {
}
}

val java17 = sourceSets.create("java17") {
java.setSrcDirs(listOf("src/main/java"))
java.include(recordSamplePath)
}

tasks.named<JavaCompile>(java17.compileJavaTaskName) {
enabled = supportsJava17
sourceCompatibility = JavaVersion.VERSION_17.toString()
targetCompatibility = JavaVersion.VERSION_17.toString()
}

tasks.jar {
if (supportsJava17) {
from(java17.output)
}
from(sourceSets.main.get().allSource) {
include("**/*.java")
}
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
package sample.alias;

public class RecordAliasSample {

public record Payload(Object value) {}

static void recordHashCodeInlined(Object src) {
Payload payload = new Payload(src);
payload.hashCode();
sinkOneValue(payload.value());
}

static void sinkOneValue(Object value) {}
}
Original file line number Diff line number Diff line change
Expand Up @@ -20,9 +20,9 @@ import org.opentaint.ir.api.jvm.JIRMethod
import org.opentaint.ir.api.jvm.JIRParameter
import org.opentaint.ir.api.jvm.JIRType
import org.opentaint.ir.api.jvm.cfg.JIRArgument
import org.opentaint.ir.api.jvm.cfg.JIRCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRImmediate
import org.opentaint.ir.api.jvm.cfg.JIRInstanceCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRMethodCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRValue
import org.opentaint.ir.api.jvm.ext.toType

Expand All @@ -33,7 +33,7 @@ sealed interface CallPositionValue {
}

class CallPositionToJIRValueResolver(
private val callExpr: JIRCallExpr,
private val callExpr: JIRMethodCallExpr,
private val returnValue: JIRImmediate?
) : PositionResolver<CallPositionValue> {
override fun resolve(position: Position): CallPositionValue = when (position) {
Expand All @@ -48,6 +48,24 @@ class CallPositionToJIRValueResolver(
}
}

class JIRMethodCallPositionBaseTypeResolver(
private val callExpr: JIRMethodCallExpr
) : PositionTypeResolver {
override fun resolve(position: PositionAccess): CommonType? {
if (position !is PositionAccess.Simple) return null

return when (val base = position.base) {
is AccessPathBase.Argument -> callExpr.args.getOrNull(base.idx)?.type
is AccessPathBase.Return -> callExpr.type
is AccessPathBase.This -> (callExpr as? JIRInstanceCallExpr)?.instance?.type
is AccessPathBase.ClassStatic,
is AccessPathBase.Constant,
is AccessPathBase.Exception,
is AccessPathBase.LocalVar -> null
}
}
}

class CalleePositionToJIRValueResolver(
private val method: JIRMethod
) : PositionResolver<CallPositionValue> {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ import org.opentaint.ir.api.jvm.cfg.JIRInst
import org.opentaint.ir.api.jvm.cfg.JIRInstanceCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRLambdaExpr
import org.opentaint.ir.api.jvm.cfg.JIRLocalVar
import org.opentaint.ir.api.jvm.cfg.JIRMethodCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRNewExpr
import org.opentaint.ir.api.jvm.cfg.JIRValue
import org.opentaint.ir.api.jvm.cfg.JIRVirtualCallExpr
Expand Down Expand Up @@ -88,18 +89,22 @@ class JIRCallResolver(
}

fun resolve(call: JIRCallExpr, location: JIRInst, context: JIRMethodAnalysisContext): List<MethodResolutionResult> {
val method = call.method.method
if (call is JIRLambdaExpr) {
// lambda expr is an allocation site. lambda calls resolved as virtual calls
return emptyList()
}

// A bootstrap method links an invokedynamic call site; it is not the
// method executed when the surrounding program reaches that instruction.
val methodCall = call as? JIRMethodCallExpr
?: return listOf(MethodResolutionResult.MethodResolutionFailed)
val method = methodCall.method.method
val methodIgnored = unitResolver.resolve(method) == UnknownUnit

if (methodIgnored && alwaysIgnoreMethod(method)) {
return listOf(MethodResolutionResult.MethodResolutionFailed)
}

if (call is JIRLambdaExpr) {
// lambda expr is an allocation site. lambda calls resolved as virtual calls
return emptyList()
}

if (call !is JIRVirtualCallExpr) {
if (methodIgnored) {
return listOf(MethodResolutionResult.MethodResolutionFailed)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import org.opentaint.ir.api.jvm.JIRClasspath
import org.opentaint.ir.api.jvm.JIRMethod
import org.opentaint.ir.api.jvm.cfg.JIRCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRInst
import org.opentaint.ir.api.jvm.cfg.JIRMethodCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRThrowInst
import org.opentaint.ir.api.jvm.cfg.JIRValue
import org.opentaint.ir.api.jvm.ext.cfg.callExpr
Expand Down Expand Up @@ -48,7 +49,8 @@ open class JIRLanguageManager(val cp: JIRClasspath) : LanguageManager {

override fun getCalleeMethod(callExpr: CommonCallExpr): JIRMethod {
jIRDowncast<JIRCallExpr>(callExpr)
return callExpr.method.method
return (callExpr as? JIRMethodCallExpr)?.method?.method
?: error("Dynamic call sites do not have a callee method")
}

override val methodContextSerializer = JIRMethodContextSerializer(cp)
Expand All @@ -60,4 +62,4 @@ internal inline fun <reified T> jIRDowncast(value: Any?) {
returns() implies(value is T)
}
check(value is T) { "Downcast error: expected ${T::class}, got $value" }
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -150,16 +150,18 @@ class DSUAliasAnalysis(

private fun evalCall(stmt: Stmt.Call, state: State, callFrame: CallTreeNode): State {
// todo: use instance alloc info
val resolvedCall = callFrame.resolveCall(stmt, methodCallResolver)
val resolvedCall = (stmt as? Stmt.MethodCall)?.let {
callFrame.resolveCall(it, methodCallResolver)
}
if (resolvedCall != null) {
val result = evalCall(stmt, state, callFrame, resolvedCall)
if (result != null) return result
}

var resultState = state
if (stmt.lValue != null) {
stmt.lValue?.let { lValue ->
val info = aliasSetFromInfo(CallReturn(stmt, callFrame.ctx))
resultState = resultState.removeOldAndMergeWith(stmt.lValue.aliasInfo().index(), info)
resultState = resultState.removeOldAndMergeWith(lValue.aliasInfo().index(), info)
}

if (!stmt.cantMutateAliasedHeap()) {
Expand All @@ -172,7 +174,8 @@ class DSUAliasAnalysis(
resultState = resultState.invalidateOuterHeapAliases(argAliases)
}

val externalModel = methodCallResolver.externalCallModel(stmt.method)
val method = (stmt as? Stmt.MethodCall)?.method
val externalModel = method?.let(methodCallResolver::externalCallModel).orEmpty()
resultState = externalModel.fold(resultState) { s, model ->
model.evalExternalCallModel(stmt, s)
}
Expand Down Expand Up @@ -566,6 +569,7 @@ class DSUAliasAnalysis(

private fun Stmt.Call.cantMutateAliasedHeap(): Boolean {
if (args.any { it !is SimpleValue.Primitive }) return false
val method = (this as? Stmt.MethodCall)?.method ?: return false
return method.isStatic || method.isConstructor
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ import org.opentaint.jvm.graph.JApplicationGraph
import java.util.BitSet

interface CallResolver {
fun resolveMethodCall(callStmt: Stmt.Call, level: Int): List<JIRMethod>?
fun resolveMethodCall(callStmt: Stmt.MethodCall, level: Int): List<JIRMethod>?
fun buildMethodGraph(method: JIRMethod): JIRInstGraph?
fun externalCallModel(method: JIRMethod): List<ExternalAssign>
}
Expand All @@ -28,7 +28,7 @@ abstract class JirCallResolver(
): CallResolver {
abstract fun buildMethodJig(entryPoint: JIRInst): JIRInstGraph

override fun resolveMethodCall(callStmt: Stmt.Call, level: Int): List<JIRMethod>? {
override fun resolveMethodCall(callStmt: Stmt.MethodCall, level: Int): List<JIRMethod>? {
if (level >= params.aliasAnalysisInterProcCallDepth) return null

val methods = callResolver.allKnownOverridesOrNull(callStmt.method)
Expand Down Expand Up @@ -56,7 +56,7 @@ class CallTreeNode(val ctx: ContextInfo, val instEvalCtx: InstEvalContext) {
private val emptyCalls = BitSet()
private val calls = Int2ObjectOpenHashMap<ResolvedCall>()

fun resolveCall(stmt: Stmt.Call, callResolver: CallResolver): Map<JIRMethod, ResolvedCallMethod>? {
fun resolveCall(stmt: Stmt.MethodCall, callResolver: CallResolver): Map<JIRMethod, ResolvedCallMethod>? {
if (emptyCalls.get(stmt.originalIdx)) return ResolvedCall.empty.methods

return calls.getOrPut(stmt.originalIdx) {
Expand Down Expand Up @@ -87,7 +87,7 @@ private class NestedCallInstEvalCtx(val call: Stmt.Call, val ctx: ContextInfo) :
override fun createLocal(idx: Int): Local = Local(idx, ctx)
}

private fun resolveCallNoCache(stmt: Stmt.Call, ctx: ContextInfo, callResolver: CallResolver): ResolvedCall {
private fun resolveCallNoCache(stmt: Stmt.MethodCall, ctx: ContextInfo, callResolver: CallResolver): ResolvedCall {
val methods = callResolver.resolveMethodCall(stmt, ctx.level)
?: return ResolvedCall.empty

Expand All @@ -106,5 +106,5 @@ private fun resolveCallNoCache(stmt: Stmt.Call, ctx: ContextInfo, callResolver:
return ResolvedCall(resolvedCall)
}

private fun mkContextId(stmt: Stmt.Call, methodIdx: Int): Int =
private fun mkContextId(stmt: Stmt.MethodCall, methodIdx: Int): Int =
(stmt.originalIdx * 1000) + methodIdx
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ import org.opentaint.ir.api.jvm.cfg.JIRImmediate
import org.opentaint.ir.api.jvm.cfg.JIRInst
import org.opentaint.ir.api.jvm.cfg.JIRInstanceCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRLocalVar
import org.opentaint.ir.api.jvm.cfg.JIRMethodCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRNewArrayExpr
import org.opentaint.ir.api.jvm.cfg.JIRNewExpr
import org.opentaint.ir.api.jvm.cfg.JIRRef
Expand Down Expand Up @@ -91,7 +92,26 @@ sealed interface Stmt : Comparable<Stmt> {

sealed interface NoCall: Stmt

data class Call(val method: JIRMethod, val lValue: RefValue.Local?, val instance: Value?, val args: List<Value>, override val originalIdx: Int) : Stmt
sealed interface Call : Stmt {
val lValue: RefValue.Local?
val instance: Value?
val args: List<Value>
}

data class MethodCall(
val method: JIRMethod,
override val lValue: RefValue.Local?,
override val instance: Value?,
override val args: List<Value>,
override val originalIdx: Int
) : Call

data class OpaqueCall(
override val lValue: RefValue.Local?,
override val instance: Value?,
override val args: List<Value>,
override val originalIdx: Int
) : Call

data class Copy(val lValue: RefValue.Local, val rValue: RefValue, override val originalIdx: Int): NoCall
data class Assign(val lValue: RefValue.Local, val expr: Expr, override val originalIdx: Int) : NoCall
Expand Down Expand Up @@ -192,15 +212,20 @@ private fun InstEvalContext.evalCall(
): Stmt? {
val lhs = (lValue as? JIRLocalVar)?.let { createLocal(it.index) }

if (expr.method.method.isPrimitiveBoxAllocMethod()) {
val method = (expr as? JIRMethodCallExpr)?.method?.method
if (method?.isPrimitiveBoxAllocMethod() == true) {
if (lhs == null) return null
return Stmt.Assign(lhs, Expr.Alloc(loc), loc.location.index)
}

val args = expr.args.map { evalSimpleValue(it as JIRImmediate, loc) }
val instance = (expr as? JIRInstanceCallExpr)?.instance?.let { evalSimpleValue(it as JIRImmediate, loc) }
val stmt = Stmt.Call(expr.method.method, lhs, instance, args, loc.location.index)
return stmt
return when (expr) {
is JIRMethodCallExpr -> Stmt.MethodCall(
expr.method.method, lhs, instance, args, loc.location.index
)
else -> Stmt.OpaqueCall(lhs, instance, args, loc.location.index)
}
}

private fun InstEvalContext.evalExpr(expr: JIRExpr, inst: JIRInst): ExprOrValue = when (expr) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ import org.opentaint.dataflow.jvm.ap.ifds.jIRDowncast
import org.opentaint.dataflow.jvm.ap.ifds.taint.JIRTaintAnalysisContext
import org.opentaint.dataflow.jvm.ap.ifds.taint.TaintRulesProvider
import org.opentaint.dataflow.jvm.ap.ifds.trace.JIRMethodCallPrecondition
import org.opentaint.dataflow.jvm.ap.ifds.trace.JIRNonMethodCallPrecondition
import org.opentaint.dataflow.jvm.ap.ifds.trace.JIRMethodSequentPrecondition
import org.opentaint.dataflow.jvm.ap.ifds.trace.JIRMethodStartPrecondition
import org.opentaint.dataflow.jvm.ifds.JIRUnitResolver
Expand All @@ -48,6 +49,7 @@ import org.opentaint.ir.api.jvm.JIRClasspath
import org.opentaint.ir.api.jvm.cfg.JIRCallExpr
import org.opentaint.ir.api.jvm.cfg.JIRImmediate
import org.opentaint.ir.api.jvm.cfg.JIRInst
import org.opentaint.ir.api.jvm.cfg.JIRMethodCallExpr
import org.opentaint.jvm.graph.JApplicationGraph
import org.opentaint.util.analysis.ApplicationGraph
import java.util.concurrent.ConcurrentHashMap
Expand Down Expand Up @@ -211,13 +213,11 @@ class JIRAnalysisManager(
jIRDowncast<JIRMethodAnalysisContext>(analysisContext)

return analysisContext.cachedCallFF(statement.location.index) {
val methodCall = callExpr as? JIRMethodCallExpr
?: return@cachedCallFF JIRNonMethodCallFlowFunction(returnValue, callExpr)

JIRMethodCallFlowFunction(
apManager,
analysisContext,
returnValue,
callExpr,
statement,
generateTrace
apManager, analysisContext, returnValue, methodCall, statement, generateTrace
)
}
}
Expand Down Expand Up @@ -259,13 +259,8 @@ class JIRAnalysisManager(
jIRDowncast<JIRInst>(statement)
jIRDowncast<JIRMethodAnalysisContext>(analysisContext)

return JIRMethodCallPrecondition(
apManager,
analysisContext,
returnValue,
callExpr,
statement
)
val methodCall = callExpr as? JIRMethodCallExpr ?: return JIRNonMethodCallPrecondition
return JIRMethodCallPrecondition(apManager, analysisContext, returnValue, methodCall, statement)
}

override fun getEdgePostProcessor(
Expand Down Expand Up @@ -327,4 +322,4 @@ class JIRAnalysisManager(
val percentValue = current.toDouble() / total
return String.format("%.2f", percentValue * 100) + "%"
}
}
}
Loading
Loading