]> git.ipfire.org Git - thirdparty/kea.git/commitdiff
Merge #1601
authorMichal 'vorner' Vaner <michal.vaner@nic.cz>
Wed, 7 Mar 2012 10:20:08 +0000 (11:20 +0100)
committerMichal 'vorner' Vaner <michal.vaner@nic.cz>
Wed, 7 Mar 2012 10:20:08 +0000 (11:20 +0100)
Conflicts:
src/bin/auth/auth_srv.cc
src/bin/auth/tests/auth_srv_unittest.cc

1  2 
src/bin/auth/auth_srv.cc
src/bin/auth/tests/auth_srv_unittest.cc

index 23751f277150aac6355aa6caaaf197d6ed87eb1f,3026f96e119f970880dbde6b8a6cf36448eb6fcf..524d581fecdde24891d64dd5c49a957a4beefee9
@@@ -291,32 -290,30 +290,32 @@@ makeErrorMessage(Message& message, Outp
      // If this is an error to a query or notify, we should also copy the
      // question section.
      if (opcode == Opcode::QUERY() || opcode == Opcode::NOTIFY()) {
-         questions.assign(message->beginQuestion(), message->endQuestion());
+         questions.assign(message.beginQuestion(), message.endQuestion());
      }
  
-     message->clear(Message::RENDER);
-     message->setQid(qid);
-     message->setOpcode(opcode);
-     message->setHeaderFlag(Message::HEADERFLAG_QR);
+     message.clear(Message::RENDER);
+     message.setQid(qid);
+     message.setOpcode(opcode);
+     message.setHeaderFlag(Message::HEADERFLAG_QR);
      if (rd) {
-         message->setHeaderFlag(Message::HEADERFLAG_RD);
+         message.setHeaderFlag(Message::HEADERFLAG_RD);
      }
      if (cd) {
-         message->setHeaderFlag(Message::HEADERFLAG_CD);
+         message.setHeaderFlag(Message::HEADERFLAG_CD);
      }
      for_each(questions.begin(), questions.end(), QuestionInserter(message));
-     message->setRcode(rcode);
+     message.setRcode(rcode);
  
 -    MessageRenderer renderer(buffer);
 +    MessageRenderer renderer;
-     renderer.setBuffer(buffer.get());
++    renderer.setBuffer(&buffer);
      if (tsig_context.get() != NULL) {
-         message->toWire(renderer, *tsig_context);
+         message.toWire(renderer, *tsig_context);
      } else {
-         message->toWire(renderer);
+         message.toWire(renderer);
      }
 +    renderer.setBuffer(NULL);
      LOG_DEBUG(auth_logger, DBG_AUTH_MESSAGES, AUTH_SEND_ERROR_RESPONSE)
-               .arg(renderer.getLength()).arg(*message);
+               .arg(renderer.getLength()).arg(message);
  }
  }
  
@@@ -481,43 -478,35 +480,43 @@@ AuthSrv::processMessage(const IOMessage
          return;
      }
  
 -    // update per opcode statistics counter.  This can only be reliable after
 -    // TSIG check succeeds.
 -    impl_->counters_.inc(message.getOpcode());
 -
      bool send_answer = true;
 -    if (message.getOpcode() == Opcode::NOTIFY()) {
 -        send_answer = impl_->processNotify(io_message, message, buffer,
 -                                           tsig_context);
 -    } else if (message.getOpcode() != Opcode::QUERY()) {
 -        LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_UNSUPPORTED_OPCODE)
 -                  .arg(message.getOpcode().toText());
 -        makeErrorMessage(message, buffer, Rcode::NOTIMP(), tsig_context);
 -    } else if (message.getRRCount(Message::SECTION_QUESTION) != 1) {
 -        makeErrorMessage(message, buffer, Rcode::FORMERR(), tsig_context);
 -    } else {
 -        ConstQuestionPtr question = *message.beginQuestion();
 -        const RRType &qtype = question->getType();
 -        if (qtype == RRType::AXFR()) {
 -            send_answer = impl_->processXfrQuery(io_message, message, buffer,
 -                                                 tsig_context);
 -        } else if (qtype == RRType::IXFR()) {
 -            send_answer = impl_->processXfrQuery(io_message, message, buffer,
 -                                                 tsig_context);
 +    try {
 +        // update per opcode statistics counter.  This can only be reliable
 +        // after TSIG check succeeds.
-         impl_->counters_.inc(message->getOpcode());
++        impl_->counters_.inc(message.getOpcode());
 +
-         if (message->getOpcode() == Opcode::NOTIFY()) {
++        if (message.getOpcode() == Opcode::NOTIFY()) {
 +            send_answer = impl_->processNotify(io_message, message, buffer,
 +                                               tsig_context);
-         } else if (message->getOpcode() != Opcode::QUERY()) {
++        } else if (message.getOpcode() != Opcode::QUERY()) {
 +            LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_UNSUPPORTED_OPCODE)
-                       .arg(message->getOpcode().toText());
++                      .arg(message.getOpcode().toText());
 +            makeErrorMessage(message, buffer, Rcode::NOTIMP(), tsig_context);
-         } else if (message->getRRCount(Message::SECTION_QUESTION) != 1) {
++        } else if (message.getRRCount(Message::SECTION_QUESTION) != 1) {
 +            makeErrorMessage(message, buffer, Rcode::FORMERR(), tsig_context);
          } else {
-             ConstQuestionPtr question = *message->beginQuestion();
 -            send_answer = impl_->processNormalQuery(io_message, message,
 -                                                    buffer, tsig_context);
++            ConstQuestionPtr question = *message.beginQuestion();
 +            const RRType &qtype = question->getType();
 +            if (qtype == RRType::AXFR()) {
 +                send_answer = impl_->processXfrQuery(io_message, message,
 +                                                     buffer, tsig_context);
 +            } else if (qtype == RRType::IXFR()) {
 +                send_answer = impl_->processXfrQuery(io_message, message,
 +                                                     buffer, tsig_context);
 +            } else {
 +                send_answer = impl_->processNormalQuery(io_message, message,
 +                                                        buffer, tsig_context);
 +            }
          }
 +    } catch (const std::exception& ex) {
 +        LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_RESPONSE_FAILURE)
 +                  .arg(ex.what());
 +        makeErrorMessage(message, buffer, Rcode::SERVFAIL());
 +    } catch (...) {
 +        LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_RESPONSE_FAILURE_UNKNOWN);
 +        makeErrorMessage(message, buffer, Rcode::SERVFAIL());
      }
 -
      impl_->resumeServer(server, message, send_answer);
  }
  
@@@ -564,19 -553,17 +563,19 @@@ AuthSrvImpl::processNormalQuery(const I
          return (true);
      }
  
 -    MessageRenderer renderer(buffer);
 +    MessageRenderer renderer;
-     renderer.setBuffer(buffer.get());
++    renderer.setBuffer(&buffer);
      const bool udp_buffer =
          (io_message.getSocket().getProtocol() == IPPROTO_UDP);
      renderer.setLengthLimit(udp_buffer ? remote_bufsize : 65535);
      if (tsig_context.get() != NULL) {
-         message->toWire(renderer, *tsig_context);
+         message.toWire(renderer, *tsig_context);
      } else {
-         message->toWire(renderer);
+         message.toWire(renderer);
      }
 +    renderer.setBuffer(NULL);
      LOG_DEBUG(auth_logger, DBG_AUTH_MESSAGES, AUTH_SEND_NORMAL_RESPONSE)
-               .arg(renderer.getLength()).arg(message->toText());
+               .arg(renderer.getLength()).arg(message);
  
      return (true);
  }
@@@ -691,18 -678,16 +690,18 @@@ AuthSrvImpl::processNotify(const IOMess
          return (false);
      }
  
-     message->makeResponse();
-     message->setHeaderFlag(Message::HEADERFLAG_AA);
-     message->setRcode(Rcode::NOERROR());
+     message.makeResponse();
+     message.setHeaderFlag(Message::HEADERFLAG_AA);
+     message.setRcode(Rcode::NOERROR());
  
 -    MessageRenderer renderer(buffer);
 +    MessageRenderer renderer;
-     renderer.setBuffer(buffer.get());
++    renderer.setBuffer(&buffer);
      if (tsig_context.get() != NULL) {
-         message->toWire(renderer, *tsig_context);
+         message.toWire(renderer, *tsig_context);
      } else {
-         message->toWire(renderer);
+         message.toWire(renderer);
      }
 +    renderer.setBuffer(NULL);
      return (true);
  }
  
index c742c5472780cb780ff7d5ffd0211d6c55473fe4,862fef5874bfb2b45cda6eed5b0682ae41aeffab..688ce62a032e163a305378e9392fbd3ba296ce53
@@@ -87,12 -87,8 +87,12 @@@ protected
          server.setXfrinSession(&notify_session);
          server.setStatisticsSession(&statistics_session);
      }
 +
      virtual void processMessage() {
-         server.processMessage(*io_message, parse_message, response_obuffer,
 +        // If processMessage has been called before, parse_message needs
 +        // to be reset. If it hasn't, there's no harm in doing so
 +        parse_message->clear(Message::PARSE);
+         server.processMessage(*io_message, *parse_message, *response_obuffer,
                                &dnsserv);
      }
  
@@@ -502,9 -492,9 +506,9 @@@ TEST_F(AuthSrvTest, AXFRDisconnectFail
                                         Name("example.com"), RRClass::IN(),
                                         RRType::AXFR());
      createRequestPacket(request_message, IPPROTO_TCP);
-     EXPECT_NO_THROW(server.processMessage(*io_message, parse_message,
-                                           response_obuffer, &dnsserv));
 -    EXPECT_THROW(server.processMessage(*io_message, *parse_message,
 -                                       *response_obuffer, &dnsserv),
 -                 XfroutError);
++    EXPECT_NO_THROW(server.processMessage(*io_message, *parse_message,
++                                          *response_obuffer, &dnsserv));
 +    // Since the disconnect failed, we should still be 'connected'
      EXPECT_TRUE(xfrout.isConnected());
      // XXX: we need to re-enable disconnect.  otherwise an exception would be
      // thrown via the destructor of the server.
@@@ -560,8 -553,9 +567,8 @@@ TEST_F(AuthSrvTest, IXFRDisconnectFail
                                         Name("example.com"), RRClass::IN(),
                                         RRType::IXFR());
      createRequestPacket(request_message, IPPROTO_TCP);
-     EXPECT_NO_THROW(server.processMessage(*io_message, parse_message,
-                                           response_obuffer, &dnsserv));
 -    EXPECT_THROW(server.processMessage(*io_message, *parse_message,
 -                                       *response_obuffer, &dnsserv),
 -                 XfroutError);
++    EXPECT_NO_THROW(server.processMessage(*io_message, *parse_message,
++                                          *response_obuffer, &dnsserv));
      EXPECT_TRUE(xfrout.isConnected());
      // XXX: we need to re-enable disconnect.  otherwise an exception would be
      // thrown via the destructor of the server.
@@@ -1050,231 -1064,4 +1075,231 @@@ TEST_F(AuthSrvTest, listenAddresses) 
                                  "Released tokens");
  }
  
-     server.processMessage(*io_message, parse_message, response_obuffer,
 +//
 +// Tests for catching exceptions in various stages of the query processing
 +//
 +// These tests work by defining two proxy classes, that act as an in-memory
 +// client by default, but can throw exceptions at various points.
 +//
 +namespace {
 +
 +/// A the possible methods to throw in, either in FakeInMemoryClient or
 +/// FakeZoneFinder
 +enum ThrowWhen {
 +    THROW_NEVER,
 +    THROW_AT_FIND_ZONE,
 +    THROW_AT_GET_ORIGIN,
 +    THROW_AT_GET_CLASS,
 +    THROW_AT_FIND,
 +    THROW_AT_FIND_ALL,
 +    THROW_AT_FIND_NSEC3
 +};
 +
 +/// convenience function to check whether and what to throw
 +void
 +checkThrow(ThrowWhen method, ThrowWhen throw_at, bool isc_exception) {
 +    if (method == throw_at) {
 +        if (isc_exception) {
 +            isc_throw(isc::Exception, "foo");
 +        } else {
 +            throw std::exception();
 +        }
 +    }
 +}
 +
 +/// \brief proxy class for the ZoneFinder returned by the InMemoryClient
 +///        proxied by FakeInMemoryClient
 +///
 +/// See the documentation for FakeInMemoryClient for more information,
 +/// all methods simply check whether they should throw, and if not, call
 +/// their proxied equivalent.
 +class FakeZoneFinder : public isc::datasrc::ZoneFinder {
 +public:
 +    FakeZoneFinder(isc::datasrc::ZoneFinderPtr zone_finder,
 +                   ThrowWhen throw_when,
 +                   bool isc_exception) :
 +        real_zone_finder_(zone_finder),
 +        throw_when_(throw_when),
 +        isc_exception_(isc_exception)
 +    {}
 +
 +    virtual isc::dns::Name
 +    getOrigin() const {
 +        checkThrow(THROW_AT_GET_ORIGIN, throw_when_, isc_exception_);
 +        return (real_zone_finder_->getOrigin());
 +    }
 +
 +    virtual isc::dns::RRClass
 +    getClass() const {
 +        checkThrow(THROW_AT_GET_CLASS, throw_when_, isc_exception_);
 +        return (real_zone_finder_->getClass());
 +    }
 +
 +    virtual isc::datasrc::ZoneFinder::FindResult
 +    find(const isc::dns::Name& name,
 +         const isc::dns::RRType& type,
 +         isc::datasrc::ZoneFinder::FindOptions options)
 +    {
 +        checkThrow(THROW_AT_FIND, throw_when_, isc_exception_);
 +        return (real_zone_finder_->find(name, type, options));
 +    }
 +
 +    virtual FindResult
 +    findAll(const isc::dns::Name& name,
 +            std::vector<isc::dns::ConstRRsetPtr> &target,
 +            const FindOptions options = FIND_DEFAULT)
 +    {
 +        checkThrow(THROW_AT_FIND_ALL, throw_when_, isc_exception_);
 +        return (real_zone_finder_->findAll(name, target, options));
 +    };
 +
 +    virtual FindNSEC3Result
 +    findNSEC3(const isc::dns::Name& name, bool recursive) {
 +        checkThrow(THROW_AT_FIND_NSEC3, throw_when_, isc_exception_);
 +        return (real_zone_finder_->findNSEC3(name, recursive));
 +    };
 +
 +    virtual isc::dns::Name
 +    findPreviousName(const isc::dns::Name& query) const {
 +        return (real_zone_finder_->findPreviousName(query));
 +    }
 +
 +private:
 +    isc::datasrc::ZoneFinderPtr real_zone_finder_;
 +    ThrowWhen throw_when_;
 +    bool isc_exception_;
 +};
 +
 +/// \brief Proxy InMemoryClient that can throw exceptions at specified times
 +///
 +/// It is based on the memory client since that one is easy to override
 +/// (with setInMemoryClient) with the current design of AuthSrv.
 +class FakeInMemoryClient : public isc::datasrc::InMemoryClient {
 +public:
 +    /// \brief Create a proxy memory client
 +    ///
 +    /// \param real_client The real in-memory client to proxy
 +    /// \param throw_when if set to any value other than never, that is
 +    ///        the method that will throw an exception (either in this
 +    ///        class or the related FakeZoneFinder)
 +    /// \param isc_exception if true, throw isc::Exception, otherwise,
 +    ///                      throw std::exception
 +    FakeInMemoryClient(AuthSrv::InMemoryClientPtr real_client,
 +                       ThrowWhen throw_when,
 +                       bool isc_exception) :
 +        real_client_(real_client),
 +        throw_when_(throw_when),
 +        isc_exception_(isc_exception)
 +    {}
 +
 +    /// \brief proxy call for findZone
 +    ///
 +    /// if this instance was constructed with throw_when set to find_zone,
 +    /// this method will throw. Otherwise, it will return a FakeZoneFinder
 +    /// instance which will throw at the method specified at the
 +    /// construction of this instance.
 +    virtual FindResult
 +    findZone(const isc::dns::Name& name) const {
 +        checkThrow(THROW_AT_FIND_ZONE, throw_when_, isc_exception_);
 +        const FindResult result = real_client_->findZone(name);
 +        return (FindResult(result.code, isc::datasrc::ZoneFinderPtr(
 +                                        new FakeZoneFinder(result.zone_finder,
 +                                        throw_when_,
 +                                        isc_exception_))));
 +    }
 +
 +private:
 +    AuthSrv::InMemoryClientPtr real_client_;
 +    ThrowWhen throw_when_;
 +    bool isc_exception_;
 +};
 +
 +} // end anonymous namespace for throwing proxy classes
 +
 +// Test for the tests
 +//
 +// Set the proxies to never throw, this should have the same result as
 +// queryWithInMemoryClientNoDNSSEC, and serves to test the two proxy classes
 +TEST_F(AuthSrvTest, queryWithInMemoryClientProxy) {
 +    // Set real inmem client to proxy
 +    updateConfig(&server, CONFIG_INMEMORY_EXAMPLE, true);
 +
 +    AuthSrv::InMemoryClientPtr fake_client(
 +        new FakeInMemoryClient(server.getInMemoryClient(rrclass),
 +                               THROW_NEVER,
 +                               false));
 +
 +    ASSERT_NE(AuthSrv::InMemoryClientPtr(), server.getInMemoryClient(rrclass));
 +    server.setInMemoryClient(rrclass, fake_client);
 +
 +    createDataFromFile("nsec3query_nodnssec_fromWire.wire");
-     server.processMessage(*io_message, parse_message, response_obuffer,
++    server.processMessage(*io_message, *parse_message, *response_obuffer,
 +                          &dnsserv);
 +
 +    EXPECT_TRUE(dnsserv.hasAnswer());
 +    headerCheck(*parse_message, default_qid, Rcode::NOERROR(),
 +                opcode.getCode(), QR_FLAG | AA_FLAG, 1, 1, 2, 1);
 +}
 +
 +// Convenience function for the rest of the tests, set up a proxy
 +// to throw in the given method
 +// If isc_exception is true, it will throw isc::Exception, otherwise
 +// it will throw std::exception
 +void
 +setupThrow(AuthSrv* server, const char *config, ThrowWhen throw_when,
 +                   bool isc_exception)
 +{
 +    // Set real inmem client to proxy
 +    updateConfig(server, config, true);
 +
 +    // Set it to throw on findZone(), this should result in
 +    // SERVFAIL on any exception
 +    AuthSrv::InMemoryClientPtr fake_client(
 +        new FakeInMemoryClient(
 +            server->getInMemoryClient(isc::dns::RRClass::IN()),
 +            throw_when,
 +            isc_exception));
 +
 +    ASSERT_NE(AuthSrv::InMemoryClientPtr(),
 +              server->getInMemoryClient(isc::dns::RRClass::IN()));
 +    server->setInMemoryClient(isc::dns::RRClass::IN(), fake_client);
 +}
 +
 +TEST_F(AuthSrvTest, queryWithThrowingProxyServfails) {
 +    // Test the common cases, all of which should simply return SERVFAIL
 +    // Use THROW_NEVER as end marker
 +    ThrowWhen throws[] = { THROW_AT_FIND_ZONE,
 +                           THROW_AT_GET_ORIGIN,
 +                           THROW_AT_FIND,
 +                           THROW_AT_FIND_NSEC3,
 +                           THROW_NEVER };
 +    UnitTestUtil::createDNSSECRequestMessage(request_message, opcode,
 +                                             default_qid, Name("foo.example."),
 +                                             RRClass::IN(), RRType::TXT());
 +    for (ThrowWhen* when(throws); *when != THROW_NEVER; ++when) {
 +        createRequestPacket(request_message, IPPROTO_UDP);
 +        setupThrow(&server, CONFIG_INMEMORY_EXAMPLE, *when, true);
 +        processAndCheckSERVFAIL();
 +        // To be sure, check same for non-isc-exceptions
 +        createRequestPacket(request_message, IPPROTO_UDP);
 +        setupThrow(&server, CONFIG_INMEMORY_EXAMPLE, *when, false);
 +        processAndCheckSERVFAIL();
 +    }
 +}
 +
 +// Throw isc::Exception in getClass(). (Currently?) getClass is not called
 +// in the processMessage path, so this should result in a normal answer
 +TEST_F(AuthSrvTest, queryWithInMemoryClientProxyGetClass) {
 +    createDataFromFile("nsec3query_nodnssec_fromWire.wire");
 +    setupThrow(&server, CONFIG_INMEMORY_EXAMPLE, THROW_AT_GET_CLASS, true);
 +
 +    // getClass is not called so it should just answer
++    server.processMessage(*io_message, *parse_message, *response_obuffer,
 +                          &dnsserv);
 +
 +    EXPECT_TRUE(dnsserv.hasAnswer());
 +    headerCheck(*parse_message, default_qid, Rcode::NOERROR(),
 +                opcode.getCode(), QR_FLAG | AA_FLAG, 1, 1, 2, 1);
 +}
 +
  }