Skip to content
Open
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
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
package org.opentaint.dataflow.ap.ifds.analysis.alias

import it.unimi.dsi.fastutil.ints.Int2ObjectOpenHashMap
import it.unimi.dsi.fastutil.longs.Long2ObjectOpenHashMap
import org.opentaint.dataflow.ap.ifds.AccessPathBase
import org.opentaint.dataflow.util.getOrCreate
Expand All @@ -15,7 +16,11 @@ abstract class LocalAliasAnalysis<AliasInfo, AliasAccessor> {

val aliasInfo: AnalysisResult? by lazy { compute() }

val convertedAliases = Long2ObjectOpenHashMap<List<AliasInfo>>()
val convertedAliases = Long2ObjectOpenHashMap<Int2ObjectOpenHashMap<List<AliasInfo>>>()

protected open val aliasCompressionThreshold: Int = Int.MAX_VALUE

protected open fun compressAliases(aliases: List<AliasInfo>): List<AliasInfo> = aliases

fun findAlias(base: AccessPathBase.LocalVar, statement: CommonInst): List<AliasInfo>? =
withStateBeforeStatement(statement) { state, stateId -> state.findLocalAlias(stateId, base.idx) }
Expand Down Expand Up @@ -120,7 +125,7 @@ abstract class LocalAliasAnalysis<AliasInfo, AliasAccessor> {
result += convert(stateId, aliasIdx, depth = 0)
}
}
return result
return compressIfRequired(result)
}

private fun State.convertAllAliasSets(stateId: Int): List<List<AliasInfo>> =
Expand All @@ -129,15 +134,19 @@ abstract class LocalAliasAnalysis<AliasInfo, AliasAccessor> {
aliasSet.forEach {
result += convert(stateId, it, depth = 0)
}
result
compressIfRequired(result)
}

abstract fun convert(info: AAInfo, depth: Int, convertInstance: (Int) -> List<AliasInfo>): List<AliasInfo>

private fun State.convert(stateId: Int, infoIdx: Int, depth: Int): List<AliasInfo> =
synchronized(convertedAliases) {
convertedAliases.getOrCreate(pair(infoIdx, stateId)) {
convert(stateId, manager.getElementUncheck(infoIdx), depth)
val cacheId = pair(infoIdx, stateId)
val resultsByDepth = convertedAliases.getOrCreate(cacheId) {
Int2ObjectOpenHashMap()
}
resultsByDepth.getOrCreate(depth) {
compressIfRequired(convert(stateId, manager.getElementUncheck(infoIdx), depth))
}
}

Expand All @@ -147,9 +156,14 @@ abstract class LocalAliasAnalysis<AliasInfo, AliasAccessor> {
forEachAliasInSet(instance) {
instances += convert(stateId, it, depth + 1)
}
instances
compressIfRequired(instances)
}

private fun compressIfRequired(aliases: List<AliasInfo>): List<AliasInfo> {
if (aliases.size <= aliasCompressionThreshold) return aliases
return compressAliases(aliases)
}

private fun pair(a: Int, b: Int): Long =
(a.toLong() shl 32) or (b.toLong() and 0xFFFF_FFFFL)
}
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,12 @@ static class Node {
Object data;
}

static class NodeExtra {
NodeExtra next;
NodeExtra prev;
Object data;
}

static void readArgField(Box box) {
Object dst = box.value;
sinkOneValue(dst);
Expand Down Expand Up @@ -90,6 +96,20 @@ static void nodeTraversalData(Node node) {
sinkOneValue(data);
}

static void twoFieldNodeTraversalData(NodeExtra node) {
NodeExtra cur = node;
while (cur.next != null && cur.prev != null) {
if (cur.data instanceof String) {
cur = cur.next;
}
else {
cur = cur.prev;
}
}
Object data = cur.data;
sinkOneValue(data);
}

static void fieldOverwrite(Box box, Object a, Object b) {
box.value = a;
box.value = b;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ import org.opentaint.dataflow.jvm.ap.ifds.alias.ExternalCallModelProvider.Extern
import org.opentaint.dataflow.jvm.ap.ifds.alias.FieldAlias
import org.opentaint.dataflow.jvm.ap.ifds.alias.JIRIntraProcAliasAnalysis
import org.opentaint.dataflow.jvm.ap.ifds.alias.JIRIntraProcAliasAnalysis.Convert.convertToAliasInfo
import org.opentaint.dataflow.jvm.ap.ifds.alias.JIRAliasPathCompressor
import org.opentaint.dataflow.jvm.ap.ifds.alias.LocalAlias
import org.opentaint.dataflow.jvm.ap.ifds.alias.RefValue
import org.opentaint.dataflow.jvm.ap.ifds.taint.TaintRulesProvider
Expand All @@ -42,6 +43,7 @@ class JIRLocalAliasAnalysis(
private val localVariableReachability: JIRLocalVariableReachability,
private val cancellation: Cancellation,
private val languageManager: JIRLanguageManager,
private val factTypeChecker: JIRFactTypeChecker,
private val params: Params,
) : LocalAliasAnalysis<AliasInfo, AliasAccessor>() {
data class Params(
Expand All @@ -66,7 +68,16 @@ class JIRLocalAliasAnalysis(
info: AAInfo,
depth: Int,
convertInstance: (Int) -> List<AliasInfo>
): List<AliasInfo> = info.convertToAliasInfo(depth, null, convertInstance)
): List<AliasInfo> = info.convertToAliasInfo(depth, null, ::isValidAccessorTransition, convertInstance)

private fun isValidAccessorTransition(previous: AliasAccessor?, next: AliasAccessor): Boolean =
isValidAliasAccessorTransition(previous, next, factTypeChecker::typeMayHaveSubtypeOf)

// Even small permutation sets multiply downstream IFDS facts, so compress every non-empty set.
override val aliasCompressionThreshold: Int = 0

override fun compressAliases(aliases: List<AliasInfo>): List<AliasInfo> =
JIRAliasPathCompressor.compress(aliases, cancellation::checkpoint)

private inner class CallModelProvider : ExternalCallModelProvider {
override fun provideModel(method: JIRMethod): List<ExternalAssign> {
Expand Down Expand Up @@ -107,10 +118,16 @@ class JIRLocalAliasAnalysis(
private fun PositionAccessor.toAaAccessor(): AAHeapAccessor? = when (this) {
is PositionAccessor.AnyFieldAccessor -> null
is PositionAccessor.ElementAccessor -> ArrayAlias
is PositionAccessor.FieldAccessor -> FieldAlias(
AliasAccessor.Field(className, fieldName, fieldType),
isImmutable = false
)
is PositionAccessor.FieldAccessor -> {
if (fieldName == "<rule-storage>") {
null
} else {
FieldAlias(
AliasAccessor.Field(className, fieldName, fieldType),
isImmutable = false
)
}
}
}
}

Expand Down Expand Up @@ -139,3 +156,17 @@ class JIRLocalAliasAnalysis(

data class AliasAllocInfo(val allocInst: Int) : AliasInfo
}

internal fun isValidAliasAccessorTransition(
previous: AliasAccessor?,
next: AliasAccessor,
typesMayOverlap: (String, String) -> Boolean,
): Boolean {
val field = next as? AliasAccessor.Field ?: return true
val previousType = when (previous) {
is AliasAccessor.Field -> previous.fieldType
is AliasAccessor.Static -> previous.typeName
is AliasAccessor.Array, null -> return true
}
return typesMayOverlap(previousType, field.className)
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,63 @@
package org.opentaint.dataflow.jvm.ap.ifds.alias

import org.opentaint.dataflow.graph.IntGraph
import org.opentaint.dataflow.jvm.ap.ifds.JIRLocalAliasAnalysis.AliasAccessor
import org.opentaint.dataflow.jvm.ap.ifds.JIRLocalAliasAnalysis.AliasApInfo
import org.opentaint.dataflow.jvm.ap.ifds.JIRLocalAliasAnalysis.AliasInfo

internal object JIRAliasPathCompressor {
fun compress(aliases: List<AliasInfo>, checkpoint: () -> Unit = {}): List<AliasInfo> {
checkpoint()
val result = LinkedHashSet(aliases)
val components = result.toList().connectedAccessorComponents(checkpoint)
if (components.isEmpty()) return result.toList()

// Dropping the whole permutation-bearing fact is intentional: a shortened path would be a new alias fact.
return result.filterNot { alias ->
checkpoint()
alias is AliasApInfo && alias.accessors.containsAccessorPermutation(components)
}
}

private fun List<AliasInfo>.connectedAccessorComponents(
checkpoint: () -> Unit,
): Map<AliasAccessor, Int> {
val accessorIds = hashMapOf<AliasAccessor, Int>()
val graph = IntGraph()

fun accessorId(accessor: AliasAccessor): Int =
accessorIds.getOrPut(accessor) { accessorIds.size }

filterIsInstance<AliasApInfo>().forEach { alias ->
checkpoint()
alias.accessors.zipWithNext { outer, inner ->
graph.addEdge(accessorId(outer), accessorId(inner))
}
}

checkpoint()
val components = graph.nonTrivialSccs().filter { it.cardinality() >= 2 }
if (components.isEmpty()) return emptyMap()

val componentByAccessorId = IntArray(accessorIds.size) { NO_COMPONENT }
components.forEachIndexed { componentId, component ->
checkpoint()
var accessorId = component.nextSetBit(0)
while (accessorId >= 0) {
componentByAccessorId[accessorId] = componentId
accessorId = component.nextSetBit(accessorId + 1)
}
}

return accessorIds.mapValues { (_, accessorId) -> componentByAccessorId[accessorId] }
}

private fun List<AliasAccessor>.containsAccessorPermutation(
componentByAccessor: Map<AliasAccessor, Int>,
): Boolean = zipWithNext().any { (outer, inner) ->
val component = componentByAccessor[outer] ?: NO_COMPONENT
component != NO_COMPONENT && component == componentByAccessor[inner]
}

private const val NO_COMPONENT = -1
}
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,7 @@ class JIRIntraProcAliasAnalysis(
fun AAInfo.convertToAliasInfo(
depth: Int,
cancellation: AnalysisCancellation?,
isValidAccessorTransition: (AliasAccessor?, AliasAccessor) -> Boolean,
resolveHeapInstance: (Int) -> List<AliasInfo>
): List<AliasInfo> {
if (this !is HeapAlias) {
Expand All @@ -118,6 +119,8 @@ class JIRIntraProcAliasAnalysis(
cancellation?.checkpoint()

val instances = resolveHeapInstance(instance)
.filterNot { it is AliasApInfo && it.accessors.size >= HEAP_CHAIN_LIMIT }

val accessor = when (val a = this.heapAccessor) {
is ArrayAlias -> AliasAccessor.Array
is FieldAlias -> a.field
Expand All @@ -127,7 +130,12 @@ class JIRIntraProcAliasAnalysis(
return instances.mapNotNull {
when (it) {
is AliasAllocInfo -> return@mapNotNull null
is AliasApInfo -> AliasApInfo(it.base, it.accessors + accessor)
is AliasApInfo -> {
if (!isValidAccessorTransition(it.accessors.lastOrNull(), accessor)) {
return@mapNotNull null
}
AliasApInfo(it.base, it.accessors + accessor)
}
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -123,7 +123,7 @@ class JIRAnalysisManager(
?: JIRLocalAliasAnalysis(
entryPointStatement, graph, callResolver.callResolver,
taintConfig,
localVariableReachability, cancellation, this, aliasAnalysisParams
localVariableReachability, cancellation, this, factTypeChecker, aliasAnalysisParams
)
} else {
null
Expand Down Expand Up @@ -328,4 +328,4 @@ class JIRAnalysisManager(
val percentValue = current.toDouble() / total
return String.format("%.2f", percentValue * 100) + "%"
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
package org.opentaint.dataflow.jvm.ap.ifds

import org.opentaint.dataflow.jvm.ap.ifds.JIRLocalAliasAnalysis.AliasAccessor
import kotlin.test.Test
import kotlin.test.assertFalse
import kotlin.test.assertTrue

class JIRAliasAccessorTypeValidationTest {
@Test
fun `rejects field access when receiver types cannot overlap`() {
val previous = field("java.io.File", "path", "java.lang.String")
val next = field("org.example.Settings", "tenantId", "org.example.TenantId")

val valid = isValidAliasAccessorTransition(previous, next) { actual, required ->
assertTrue(actual == "java.lang.String")
assertTrue(required == "org.example.Settings")
false
}

assertFalse(valid)
}

@Test
fun `accepts field access when receiver types may overlap`() {
val previous = field("org.example.Container", "value", "org.example.HasTenant")
val next = field("org.example.Settings", "tenantId", "org.example.TenantId")

assertTrue(isValidAliasAccessorTransition(previous, next) { _, _ -> true })
}

@Test
fun `validates a field following a static base`() {
val previous = AliasAccessor.Static("org.example.Settings")
val next = field("org.example.Settings", "DEFAULT", "org.example.Settings")

assertTrue(isValidAliasAccessorTransition(previous, next) { actual, required ->
actual == required
})
}

@Test
fun `keeps transitions without enough type information`() {
val field = field("org.example.Settings", "tenantId", "org.example.TenantId")
val rejectAll: (String, String) -> Boolean = { _, _ -> false }

assertTrue(isValidAliasAccessorTransition(null, field, rejectAll))
assertTrue(isValidAliasAccessorTransition(AliasAccessor.Array, field, rejectAll))
assertTrue(isValidAliasAccessorTransition(field, AliasAccessor.Array, rejectAll))
}

private fun field(className: String, fieldName: String, fieldType: String) =
AliasAccessor.Field(className, fieldName, fieldType)
}
Original file line number Diff line number Diff line change
Expand Up @@ -452,6 +452,33 @@ class AliasSampleTest : BasicTestUtils() {
}
}

@Test
fun `test node traversal on two fields produces field chain ending with data`() {
val method = findMethod(HEAP_SAMPLE, "twoFieldNodeTraversalData")
val aa = aaForMethod(method)

val sink = method.findSinkCall("sinkOneValue")
val apAliases = aa.sinkArgApAliases(sink)

assertTrue {
apAliases.any {
it.base == Argument(0)
&& it.accessors.size >= 2
&& it.accessors.last().isField(FIELD_DATA)
&& it.accessors.dropLast(1).all { a -> a.isField(FIELD_NEXT) }
}
}

assertTrue {
apAliases.any {
it.base == Argument(0)
&& it.accessors.size >= 2
&& it.accessors.last().isField(FIELD_DATA)
&& it.accessors.dropLast(1).all { a -> a.isField(FIELD_PREV) }
}
}
}

@Test
fun `test field overwrite on argument receiver`() {
val method = findMethod(HEAP_SAMPLE, "fieldOverwrite")
Expand Down Expand Up @@ -589,7 +616,10 @@ class AliasSampleTest : BasicTestUtils() {
val localReachability = JIRLocalVariableReachability(method, graph, manager)
val cancellation = Cancellation().also { it.activate() }

return JIRLocalAliasAnalysis(ep, graph, callResolver, noRules, localReachability, cancellation, manager, params)
return JIRLocalAliasAnalysis(
ep, graph, callResolver, noRules,
localReachability, cancellation, manager, manager.factTypeChecker, params
)
}

private fun interProcParams(depth: Int) =
Expand Down Expand Up @@ -643,6 +673,7 @@ class AliasSampleTest : BasicTestUtils() {
private const val FIELD_VALUE = "value"
private const val FIELD_BOX = "box"
private const val FIELD_NEXT = "next"
private const val FIELD_PREV = "prev"
private const val FIELD_DATA = "data"
private const val FIELD_INTERPROC = "field"
}
Expand Down
Loading
Loading