Skip to content

Commit bf97136

Browse files
committed
Allow functions with multiple return values in Guards logic
1 parent 7f129bc commit bf97136

2 files changed

Lines changed: 213 additions & 29 deletions

File tree

  • go/ql
    • lib/semmle/go/controlflow
    • test/library-tests/semmle/go/dataflow/GuardingFunctions

‎go/ql/lib/semmle/go/controlflow/Guards.qll‎

Lines changed: 80 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -303,33 +303,80 @@ private module GuardsInput implements
303303
pragma[inline]
304304
predicate parameterMatch(ParameterPosition ppos, ArgumentPosition apos) { ppos = apos }
305305

306-
final private class FinalFunction = G::Function;
306+
private newtype TNonOverridableMethod =
307+
TMethod(G::Function function, int resultIndex) {
308+
exists(function.getFuncDecl()) and
309+
(
310+
function.getNumResult() = 0 and resultIndex = -1
311+
or
312+
resultIndex in [0 .. function.getNumResult() - 1]
313+
)
314+
}
307315

308316
/**
309-
* A declared function or concrete method.
317+
* A result of a declared function or concrete method, or its normal
318+
* completion when it has no results.
310319
*
311320
* Calls are restricted separately to calls whose syntactic target is this
312321
* function or method, excluding interface dispatch.
313322
*/
314-
class NonOverridableMethod extends FinalFunction {
315-
NonOverridableMethod() {
316-
exists(super.getFuncDecl()) and
317-
super.getNumResult() <= 1
323+
class NonOverridableMethod extends TNonOverridableMethod {
324+
G::Function getFunction() { this = TMethod(result, _) }
325+
326+
int getResultIndex() { this = TMethod(_, result) }
327+
328+
string toString() {
329+
result = this.getFunction().toString() + " result " + this.getResultIndex().toString()
318330
}
319331

320-
Parameter getParameter(ParameterPosition ppos) { result = super.getParameter(ppos) }
332+
int getNumParameter() { result = this.getFunction().getNumParameter() }
333+
334+
Parameter getParameter(ParameterPosition ppos) {
335+
result = this.getFunction().getParameter(ppos)
336+
}
337+
338+
/**
339+
* Holds if every return maps one expression to each result position.
340+
*
341+
* Otherwise, the shared wrapper analysis would treat a partial set of
342+
* return expressions as exhaustive.
343+
*/
344+
private predicate hasOnlyPositionMappedReturns() {
345+
forall(G::ReturnStmt ret | ret.getEnclosingFunction() = this.getFunction().getFuncDecl() |
346+
ret.getNumExpr() = this.getFunction().getNumResult()
347+
)
348+
}
321349

322350
/** Gets an expression being returned by this function. */
323351
Expr getAReturnExpr() {
324352
exists(G::ReturnStmt ret |
325-
ret.getEnclosingFunction() = super.getFuncDecl() and
326-
result = ret.getExpr()
353+
this.getResultIndex() >= 0 and
354+
this.hasOnlyPositionMappedReturns() and
355+
ret.getEnclosingFunction() = this.getFunction().getFuncDecl() and
356+
result = ret.getExpr(this.getResultIndex())
327357
)
328358
}
329359
}
330360

331-
private predicate nonOverridableCall(G::CallExpr call, NonOverridableMethod m) {
332-
call.getTarget() = m
361+
private predicate extractedCallResult(Expr use, G::CallExpr call, int resultIndex) {
362+
exists(GoSsa::SsaDefinition def, IR::ExtractTupleElementInstruction extract |
363+
use = def.getVariable().getAUse().(IR::EvalInstruction).getExpr() and
364+
def.(GoSsa::SsaExplicitDefinition).getInstruction() = extract and
365+
extract.extractsElement(IR::evalExprInstruction(call), resultIndex)
366+
)
367+
}
368+
369+
private predicate nonOverridableCall(
370+
Expr resultExpr, G::CallExpr call, NonOverridableMethod method
371+
) {
372+
call.getTarget() = method.getFunction() and
373+
(
374+
method.getFunction().getNumResult() <= 1 and
375+
resultExpr = call
376+
or
377+
method.getFunction().getNumResult() > 1 and
378+
extractedCallResult(resultExpr, call, method.getResultIndex())
379+
)
333380
}
334381

335382
private predicate hasExplicitReceiverArgument(G::CallExpr call) {
@@ -351,28 +398,32 @@ private module GuardsInput implements
351398
)
352399
}
353400

354-
class NonOverridableMethodCall extends Expr instanceof G::CallExpr {
355-
NonOverridableMethodCall() { nonOverridableCall(this, _) }
401+
class NonOverridableMethodCall extends Expr {
402+
NonOverridableMethodCall() { nonOverridableCall(this, _, _) }
403+
404+
private G::CallExpr getCall() { nonOverridableCall(this, result, _) }
356405

357-
NonOverridableMethod getMethod() { nonOverridableCall(this, result) }
406+
NonOverridableMethod getMethod() { nonOverridableCall(this, _, result) }
358407

359408
Expr getArgument(ArgumentPosition apos) {
360-
(
361-
not hasExplicitReceiverArgument(this) and
409+
exists(G::CallExpr call | call = this.getCall() |
362410
(
363-
apos = -1 and
364-
result = getDirectReceiverArgument(this, this.getMethod())
411+
not hasExplicitReceiverArgument(call) and
412+
(
413+
apos = -1 and
414+
result = getDirectReceiverArgument(call, this.getMethod())
415+
or
416+
apos != -1 and
417+
result = call.getArgument(apos)
418+
)
365419
or
366-
apos != -1 and
367-
result = super.getArgument(apos)
420+
hasExplicitReceiverArgument(call) and
421+
result = call.getArgument(apos + 1)
422+
) and
423+
not (
424+
call.hasImplicitVarargs() and
425+
apos = this.getMethod().getNumParameter() - 1
368426
)
369-
or
370-
hasExplicitReceiverArgument(this) and
371-
result = super.getArgument(apos + 1)
372-
) and
373-
not (
374-
super.hasImplicitVarargs() and
375-
apos = this.getMethod().getNumParameter() - 1
376427
)
377428
}
378429
}
@@ -412,8 +463,8 @@ private module LogicInput implements GuardsImpl::LogicInputSig {
412463

413464
predicate implicitReturnDefinition(GuardsInput::NonOverridableMethod method, SsaDefinition def) {
414465
exists(IR::ReadResultInstruction read |
415-
method.getNumResult() = 1 and
416-
read.reads(method.getResult(0)) and
466+
method.getResultIndex() >= 0 and
467+
read.reads(method.getFunction().getResult(method.getResultIndex())) and
417468
def.getVariable().getAUse() = read
418469
)
419470
}

‎go/ql/test/library-tests/semmle/go/dataflow/GuardingFunctions/test.go‎

Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -334,6 +334,62 @@ func deeplyNestedConditionalRight(p string) bool {
334334
return p[1] == 'b' && len(p)%2 == 1 && p[0] == 'a' && !isBad(p)
335335
}
336336

337+
// Valid when the second result is nil
338+
func guardMultiError(p string) (string, error) {
339+
if isBad(p) {
340+
return "", errors.New("invalid")
341+
}
342+
return p, nil
343+
}
344+
345+
// Valid when the first result is true; the second result is unrelated
346+
func guardMultiResultIsolation(p string) (bool, bool) {
347+
return !isBad(p), true
348+
}
349+
350+
// Valid when the named error result is nil
351+
func guardMultiNamed(p string) (value string, err error) {
352+
if isBad(p) {
353+
err = errors.New("invalid")
354+
return
355+
}
356+
value = p
357+
return
358+
}
359+
360+
// Not a guard: the naked return can return false without validating p
361+
func mixedNamedResultGuard(p string, bypass bool) (invalid bool) {
362+
if bypass {
363+
return
364+
}
365+
return isBad(p)
366+
}
367+
368+
func uncheckedMultiResult(p string) (string, error) {
369+
return p, nil
370+
}
371+
372+
// Not a guard: the tuple-forwarding return can return nil without validating p
373+
func mixedTupleReturnGuard(p string, bypass bool) (string, error) {
374+
if bypass {
375+
return uncheckedMultiResult(p)
376+
}
377+
if isBad(p) {
378+
return "", errors.New("invalid")
379+
}
380+
return p, nil
381+
}
382+
383+
type multiGuard struct{}
384+
385+
// Validates p when the second result is nil
386+
func (multiGuard) validate(p string) (string, error) {
387+
if isBad(p) {
388+
return "", errors.New("invalid")
389+
}
390+
return p, nil
391+
}
392+
337393
// Finally, actually test the functions -- try sinking a tainted value in the is-true/false
338394
// or is-nil/non-nil case for each candidate:
339395

@@ -848,4 +904,81 @@ func test() {
848904
}
849905
}
850906

907+
{
908+
s := source()
909+
_, err := guardMultiError(s)
910+
if err == nil {
911+
sink(s)
912+
} else {
913+
sink(s) // $ hasValueFlow="s"
914+
}
915+
}
916+
917+
{
918+
s := source()
919+
valid, _ := guardMultiResultIsolation(s)
920+
if valid {
921+
sink(s)
922+
} else {
923+
sink(s) // $ hasValueFlow="s"
924+
}
925+
}
926+
927+
{
928+
s := source()
929+
_, unrelated := guardMultiResultIsolation(s)
930+
if unrelated {
931+
sink(s) // $ hasValueFlow="s"
932+
} else {
933+
sink(s) // $ hasValueFlow="s"
934+
}
935+
}
936+
937+
{
938+
s := source()
939+
_, err := guardMultiNamed(s)
940+
if err == nil {
941+
sink(s)
942+
} else {
943+
sink(s) // $ hasValueFlow="s"
944+
}
945+
}
946+
947+
{
948+
s := source()
949+
invalid := mixedNamedResultGuard(s, true)
950+
if !invalid {
951+
sink(s) // $ hasValueFlow="s"
952+
}
953+
}
954+
955+
{
956+
s := source()
957+
_, err := mixedTupleReturnGuard(s, true)
958+
if err == nil {
959+
sink(s) // $ hasValueFlow="s"
960+
}
961+
}
962+
963+
{
964+
s := source()
965+
_, err := guardMultiError(s)
966+
copiedErr := err
967+
if copiedErr == nil {
968+
sink(s)
969+
} else {
970+
sink(s) // $ hasValueFlow="s"
971+
}
972+
}
973+
974+
{
975+
s := source()
976+
_, err := (multiGuard{}).validate(s)
977+
if err == nil {
978+
sink(s)
979+
} else {
980+
sink(s) // $ hasValueFlow="s"
981+
}
982+
}
983+
851984
}

0 commit comments

Comments
 (0)