Skip to content

Commit 97e898b

Browse files
committed
Allow functions with multiple return values in Guards logic
1 parent 857f54e commit 97e898b

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
@@ -304,33 +304,80 @@ private module GuardsInput implements
304304
pragma[inline]
305305
predicate parameterMatch(ParameterPosition ppos, ArgumentPosition apos) { ppos = apos }
306306

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

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

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

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

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

336383
private predicate hasExplicitReceiverArgument(G::CallExpr call) {
@@ -352,28 +399,32 @@ private module GuardsInput implements
352399
)
353400
}
354401

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

358-
NonOverridableMethod getMethod() { nonOverridableCall(this, result) }
407+
NonOverridableMethod getMethod() { nonOverridableCall(this, _, result) }
359408

360409
Expr getArgument(ArgumentPosition apos) {
361-
(
362-
not hasExplicitReceiverArgument(this) and
410+
exists(G::CallExpr call | call = this.getCall() |
363411
(
364-
apos = -1 and
365-
result = getDirectReceiverArgument(this, this.getMethod())
412+
not hasExplicitReceiverArgument(call) and
413+
(
414+
apos = -1 and
415+
result = getDirectReceiverArgument(call, this.getMethod())
416+
or
417+
apos != -1 and
418+
result = call.getArgument(apos)
419+
)
366420
or
367-
apos != -1 and
368-
result = super.getArgument(apos)
421+
hasExplicitReceiverArgument(call) and
422+
result = call.getArgument(apos + 1)
423+
) and
424+
not (
425+
call.hasImplicitVarargs() and
426+
apos = this.getMethod().getNumParameter() - 1
369427
)
370-
or
371-
hasExplicitReceiverArgument(this) and
372-
result = super.getArgument(apos + 1)
373-
) and
374-
not (
375-
super.hasImplicitVarargs() and
376-
apos = this.getMethod().getNumParameter() - 1
377428
)
378429
}
379430
}
@@ -413,8 +464,8 @@ private module LogicInput implements GuardsImpl::LogicInputSig {
413464

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

‎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)